MAX Pipeline 数据预处理深入解析:Batch Padding、因果注意力掩码与生成长度控制
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
导读
在 MAX(Modular Accelerated Xecution)平台的 Python 推理流水线中,max.pipelines.modeling.dataprocessing模块负责将不同长度的输入 token 序列整理为统一形状的批量张量、为变长序列构造因果注意力掩码,并计算每轮生成的最大 token 数。这三个环节直接决定了批处理效率与解码质量。本文将结合该模块的源码实现与集成测试,逐函数讲解其设计意图、参数语义、调用关系与典型应用场景,帮助你在自定义架构或流水线集成中正确使用这些预处理工具。
模块定位与整体结构
max.pipelines.modeling.dataprocessing位于仓库的 max/python/max/pipelines/modeling/dataprocessing/ 目录,共包含四个源文件与一个包入口:
- init.py:导出全部公开 API;
- collate_batch.py:批量拼接(Batch collation)逻辑,包含
PaddingDirection、collate_batch、batch_padded_tokens_and_mask; - causal_attention_mask.py:因果注意力掩码构造,包含
causal_attention_mask与causal_attention_mask_with_token_mask; - max_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 = "right"LEFT:在序列左侧(头部)填充,适合解码阶段——KV Cache 中位置对齐需要每个样本的"最后一个 token"都落在同一列,左侧填充可保持末尾位置一致;RIGHT:在序列右侧(尾部)填充,默认值,适合编码阶段(如 BERT 类模型),填充位不会干扰左侧的真实 token 顺序。
collate_batch:核心参数与行为
函数签名如下(collate_batch.py):
def 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]]]:| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
batch | list[np.ndarray[int64]] | 必填 | 一批一维 int64 token 数组,长度可不同 |
direction | PaddingDirection | RIGHT | 填充方向 |
pad_value | int | 0 | 填充使用的 token id,实践中通常传 pad token id |
batch_size | int \| None | None | 目标批量大小;传入时用全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_len; LEFT方向下,每个样本的"末 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),mode="constant"表示用固定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_size=len(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.py):
def 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, k=start_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 加上当前及之前的新 token;post_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_mask:
def 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 或需要隐藏的 token;rank-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_VAL(L122-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_inputs(batch_processor.py)从self.runtime.pad_token_id取得填充 id(_pad_token_id,见 L803-L805),对每个TextContext的激活 token 调用collate_batch(tokens, pad_value=pad_value, batch_size=len(tokens))得到固定形状的批量张量,随后以next_tokens_batch != pad_value生成 float32 注意力掩码并转为设备端Buffer。这展示了pad_value参数在真实流水线中取 pad token id 的用法。多模态文本编码器掩码构造:Qwen3 文本编码器的
attention_bias_from_attention_mask_array(qwen3/text_encoder/model.py)调用causal_attention_mask_with_token_mask([0], attention_mask)构造加性掩码,校验 batch=1 与序列长度后,通过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 testing(hypothesis 随机生成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),仅供参考