
MAX Pipeline 数据预处理深入解析Batch Padding、因果注意力掩码与生成长度控制【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo导读在 MAXModular Accelerated Xecution平台的 Python 推理流水线中max.pipelines.modeling.dataprocessing模块负责将不同长度的输入 token 序列整理为统一形状的批量张量、为变长序列构造因果注意力掩码并计算每轮生成的最大 token 数。这三个环节直接决定了批处理效率与解码质量。本文将结合该模块的源码实现与集成测试逐函数讲解其设计意图、参数语义、调用关系与典型应用场景帮助你在自定义架构或流水线集成中正确使用这些预处理工具。模块定位与整体结构max.pipelines.modeling.dataprocessing位于仓库的 max/python/max/pipelines/modeling/dataprocessing/ 目录共包含四个源文件与一个包入口init.py导出全部公开 APIcollate_batch.py批量拼接Batch collation逻辑包含PaddingDirection、collate_batch、batch_padded_tokens_and_maskcausal_attention_mask.py因果注意力掩码构造包含causal_attention_mask与causal_attention_mask_with_token_maskmax_tokens_to_generate.py生成长度计算工具。从模块导出表__all__可见公开 API 为 6 个符号PaddingDirection、batch_padded_tokens_and_mask、causal_attention_mask、causal_attention_mask_with_token_mask、collate_batch、max_tokens_to_generate与 Sphinx API 文档 pipelines.modeling.dataprocessing.rst 中按 Batch collation、Attention masks、Utilities 三组列出的内容一一对应。文档中使用的.. automodule::/.. autosummary::指令表明该页为自动生成的模块级 API 参考具体行为以源码与测试为准。Batch collation把变长序列拼成批量张量推理引擎一次处理一批请求而每个请求的 prompt 长度不同。collate_batch的核心任务就是把这一批长度不一的 token 序列统一填充pad到相同的长度并记录每个样本的真实末 token 位置供后续取 logits 使用。PaddingDirection填充方向枚举collate_batch.py 定义了填充方向枚举class PaddingDirection(enum.Enum): Padding (from) direction for batch collation. LEFT left RIGHT rightLEFT在序列左侧头部填充适合解码阶段——KV Cache 中位置对齐需要每个样本的最后一个 token都落在同一列左侧填充可保持末尾位置一致RIGHT在序列右侧尾部填充默认值适合编码阶段如 BERT 类模型填充位不会干扰左侧的真实 token 顺序。collate_batch核心参数与行为函数签名如下collate_batch.pydef collate_batch( batch: list[npt.NDArray[np.int64]], direction: PaddingDirection PaddingDirection.RIGHT, pad_value: int 0, batch_size: int | None None, ) - tuple[npt.NDArray[np.integer[Any]], npt.NDArray[np.integer[Any]]]:参数类型默认值说明batchlist[np.ndarray[int64]]必填一批一维 int64 token 数组长度可不同directionPaddingDirectionRIGHT填充方向pad_valueint0填充使用的 token id实践中通常传 pad token idbatch_sizeint \| NoneNone目标批量大小传入时用全pad_value的占位序列补齐到该大小返回值为二元组形状为(batch_size, max_seq_len)的填充后矩阵所有行长度一致长度为batch_size的未填充末 token 索引数组。具体行为与约束依据源码 collate_batch.py空 batch 直接抛出ValueError(Must provide at least one batch item.)目前仅支持一维rank-1张量若输入包含高维数组会抛出NotImplementedError(Collate only supports rank 1 tensors for now.)填充长度取 batch 内最大序列长度max_lenLEFT方向下每个样本的末 token 索引统一为-1即填充后矩阵的最后一列RIGHT方向下则为各样本原始长度减一len(a) - 1若指定了batch_size且大于当前样本数会用长度为pad_to、值全为pad_value的占位行补齐这一机制在流水线中用于将 batch 扩充到设备要求的固定形状。np.pad的填充实现collate_batch.py用npad计算差值LEFT时在头部补(npad, 0)RIGHT时在尾部补(0, npad)modeconstant表示用固定pad_value填充。batch_padded_tokens_and_mask一站式打包collate_batch.py 提供的组合函数把 token 填充与掩码构造合并def batch_padded_tokens_and_mask( start_pos: list[int], tokens: list[npt.NDArray[np.int64]], ) - tuple[npt.NDArray[np.integer[Any]], npt.NDArray[np.integer[Any]], npt.NDArray[np.float32]]:参数语义start_pos每个 batch 样本在 KV Cache 中的起始位置即该样本此前已编码的 context 长度tokens本轮待处理的未填充 token 序列列表。它内部先以start_pos与各序列长度调用causal_attention_mask生成注意力掩码再调用collate_batch(tokens, batch_sizelen(tokens))完成填充返回三元组(填充后的 token 批量张量, 未填充末 token 索引, 注意力掩码)。集成测试 test_collate_batch.py 验证了返回三者的形状约束batched_tokens.shape[0] len(tokens)、末 token 索引长度与 batch 一致、attention_mask.shape[:2] batched_tokens.shape。Attention masks为变长序列构造因果掩码因果注意力要求每个 token 只能看到它自己及之前的 token。模块用**加性掩码additive mask**实现可见位置为0不可见位置为大负数_FILL_VAL -10000.0掩码被加到注意力分数上再做 softmax从而把不可见位置的注意力权重压到接近 0。为什么用 -10000.0 而不是 -inf源码 causal_attention_mask.py 与测试文件 test_causal_attention_mask.py 都保留了同一注释TODO(KERN-782): This should be -inf but softmax saturates with NaNs.即数学上应使用-inf但 softmax 在极端值下会饱和并产生 NaN因此工程上采用-10000.0这一足够大且数值安全的负值。这是从源码可确认的实现事实也解释了测试中对FILL_VAL的固定断言。causal_attention_mask带 context 偏移的因果掩码签名causal_attention_mask.pydef causal_attention_mask( original_start_pos: list[int], original_seq_len: list[int], ) - npt.NDArray[np.float32]:构造逻辑causal_attention_mask.py将start_pos与seq_len转为 int64 数组padded_length seq_len.max()作为本批统一的新增 token 数计算post_seq_len (start_pos padded_length).max()即本轮结束后最长的总上下文长度生成形状(padded_length, post_seq_len)的_FILL_VAL填充矩阵对每个样本以np.triu(fill_matrix, kstart_pos 1)取严格上三角k start_pos 1的作用是让 token 能 attend 到自身源码注释 Set diagonal to k 1 so that tokens attend to themselves将所有样本的掩码np.stack成形状(batch, padded_length, post_seq_len)的 float32 数组。换言之第i个样本的掩码中第pos行的可见区间为[0, start_pos pos 1)即此前全部 context 加上当前及之前的新 tokenpost_seq_len之后的列padding 区域全部为_FILL_VAL。集成测试分别验证了形状test_causal_attention_mask.py、padding 被屏蔽L69-L77、当前与后续 token 被屏蔽L85-L93、以及先前 token 可见性L101-L108四个性质。causal_attention_mask_with_token_mask叠加 token 级有效性掩码某些场景如多模态输入、padding 语义变化需要在因果掩码之上再屏蔽无效 token。该函数causal_attention_mask.py在因果掩码基础上叠加一层token_maskdef causal_attention_mask_with_token_mask( original_start_pos: list[int], token_mask: npt.ArrayLike, *, mask_name: str token_mask, ) - npt.NDArray[np.float32]:关键行为token_mask支持 rank-1[seq_len]或 rank-2[batch, seq_len]的 bool/类 bool 数组True表示有效、False表示 padding 或需要隐藏的 tokenrank-1 会被自动扩展为 rank-2非法维度非 1 维、非 2 维抛出ValueError错误消息中会带上mask_name默认token_mask便于在调用方定位参数要求len(original_start_pos) batch_size否则抛出ValueError先调用causal_attention_mask得到基础加性掩码再把token_mask按各样本start_pos对齐放入总上下文区间causal_attention_mask.py最后用np.where把无效位置的掩码值替换为_FILL_VALL122-L126。测试 test_causal_attention_mask.py 给出的实例直观展示了叠加效果对start_pos[0]、token_mask[True, False, True, False]输出掩码中第 1、3 列False位置整列为FILL_VAL其余按因果规则保留0.0。第二个测试L130-L146验证了start_pos[2]时前缀 context 保持可见前三列全为0.0同时无效 token 仍被屏蔽。Utilities生成长度上限计算max_tokens_to_generate.py 提供max_tokens_to_generate工具def max_tokens_to_generate( prompt_size: int, max_length: int, max_new_tokens: int -1, ) - int:计算规则difference max(max_length - prompt_size, 0)总长度上限减去已消耗的 prompt 长度下限钳制为 0若max_new_tokens 0默认-1只受max_length约束返回difference否则返回min(max_new_tokens, difference)即同时尊重两个上限、取更严格者。该函数在仓库内还被 max/python/max/pipelines/lib/tokenizer.py 与 max/python/max/pipelines/architectures/idefics3/tokenizer.py 等 tokenizer 工具中复用用于在请求层面统一计算本轮允许生成的新 token 数。真实调用链从流水线到内核上述工具并非孤立存在而是嵌入 MAX pipeline 的实际执行路径中以下调用点均可在仓库中直接验证编码器批量输入准备PaddedEncoderBatchProcessor.prepare_initial_token_inputsbatch_processor.py从self.runtime.pad_token_id取得填充 id_pad_token_id见 L803-L805对每个TextContext的激活 token 调用collate_batch(tokens, pad_valuepad_value, batch_sizelen(tokens))得到固定形状的批量张量随后以next_tokens_batch ! pad_value生成 float32 注意力掩码并转为设备端Buffer。这展示了pad_value参数在真实流水线中取 pad token id 的用法。多模态文本编码器掩码构造Qwen3 文本编码器的attention_bias_from_attention_mask_arrayqwen3/text_encoder/model.py调用causal_attention_mask_with_token_mask([0], attention_mask)构造加性掩码校验 batch1 与序列长度后通过additive_mask[:, np.newaxis, :, :]扩展为 4 维 attention bias 供注意力核使用。causal_attention_mask_with_token_mask同样被 qwen3_modulev3/text_encoder/model.py 引用说明该工具在多个架构中通用。组合使用batch_padded_tokens_and_mask将填充与掩码构造串成一步其返回值形状约束由 test_collate_batch.py 以 property-based testinghypothesis 随机生成start_pos与tokens保证。设计要点与使用建议综合源码与测试可提炼出以下设计结论其中推断性表述均以源码结构为依据填充方向与解码阶段的配合LEFT填充把真实末 token 统一到最后一列索引-1配合 KV Cache 位置对齐RIGHT填充保持从左到右的自然顺序适合编码器。选择方向时需与后续 logits 提取逻辑一致返回的unpadded_last_token_index正是为此设计。加性掩码与数值安全掩码统一使用-10000.0而非-inf是避免 softmax NaN 的工程取舍这是源码注释明确记录的 TODO 事项KERN-782集成时应沿用同一约定。形状契约清晰causal_attention_mask输出形状为(batch, padded_length, post_seq_len)其中padded_length是 batch 内最大序列长度、post_seq_len是加 context 后的最大总长度调用方必须按此约定组织后续注意力计算。校验完整、错误可定位空 batch、高维输入、token_mask维度错误、start_pos与 batch 大小不匹配等边界条件均有显式异常与可定制错误名mask_name便于在复杂流水线中快速定位。生成长度的双上限取小语义max_tokens_to_generate把总长度上限与新增 token 上限统一为最小化约束且负值max_new_tokens表示不设新增上限这是请求级解码控制的基础。如需深入可直接阅读模块源码 collate_batch.py、causal_attention_mask.py、max_tokens_to_generate.py以及集成测试 test_collate_batch.py 与 test_causal_attention_mask.py其中包含了完整的随机属性测试与数值断言是理解各函数语义边界的权威参考。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考