news 2026/9/12 14:42:06

MAX Pipeline 数据预处理深入解析:Batch Padding、因果注意力掩码与生成长度控制

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MAX Pipeline 数据预处理深入解析:Batch Padding、因果注意力掩码与生成长度控制

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)逻辑,包含PaddingDirectioncollate_batchbatch_padded_tokens_and_mask
  • causal_attention_mask.py:因果注意力掩码构造,包含causal_attention_maskcausal_attention_mask_with_token_mask
  • max_tokens_to_generate.py:生成长度计算工具。

从模块导出表(__all__)可见,公开 API 为 6 个符号:PaddingDirectionbatch_padded_tokens_and_maskcausal_attention_maskcausal_attention_mask_with_token_maskcollate_batchmax_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]]]:
参数类型默认值说明
batchlist[np.ndarray[int64]]必填一批一维 int64 token 数组,长度可不同
directionPaddingDirectionRIGHT填充方向
pad_valueint0填充使用的 token id,实践中通常传 pad token id
batch_sizeint \| NoneNone目标批量大小;传入时用全pad_value的占位序列补齐到该大小

返回值为二元组:

  1. 形状为(batch_size, max_seq_len)的填充后矩阵,所有行长度一致;
  2. 长度为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):

  1. start_posseq_len转为 int64 数组,padded_length = seq_len.max()作为本批统一的新增 token 数;
  2. 计算post_seq_len = (start_pos + padded_length).max(),即本轮结束后最长的总上下文长度;
  3. 生成形状(padded_length, post_seq_len)_FILL_VAL填充矩阵;
  4. 对每个样本,以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");
  5. 将所有样本的掩码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:

计算规则:

  1. difference = max(max_length - prompt_size, 0):总长度上限减去已消耗的 prompt 长度,下限钳制为 0;
  2. max_new_tokens < 0(默认-1),只受max_length约束,返回difference
  3. 否则返回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_postokens)保证。

设计要点与使用建议

综合源码与测试,可提炼出以下设计结论(其中推断性表述均以源码结构为依据):

  1. 填充方向与解码阶段的配合LEFT填充把真实末 token 统一到最后一列(索引-1),配合 KV Cache 位置对齐;RIGHT填充保持从左到右的自然顺序,适合编码器。选择方向时需与后续 logits 提取逻辑一致(返回的unpadded_last_token_index正是为此设计)。
  2. 加性掩码与数值安全:掩码统一使用-10000.0而非-inf是避免 softmax NaN 的工程取舍,这是源码注释明确记录的 TODO 事项(KERN-782),集成时应沿用同一约定。
  3. 形状契约清晰causal_attention_mask输出形状为(batch, padded_length, post_seq_len),其中padded_length是 batch 内最大序列长度、post_seq_len是加 context 后的最大总长度;调用方必须按此约定组织后续注意力计算。
  4. 校验完整、错误可定位:空 batch、高维输入、token_mask维度错误、start_pos与 batch 大小不匹配等边界条件均有显式异常与可定制错误名(mask_name),便于在复杂流水线中快速定位。
  5. 生成长度的"双上限取小"语义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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 14:40:08

旧纺织品回收分类与变现全攻略

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 14:39:40

Java校园卡系统实战:Eclipse+HSQLDB+SWT+JSP轻量闭环开发

简介&#xff1a;这是一份基于Java开发的轻量级校园卡管理系统实战项目&#xff0c;面向Java初学者与课程设计学生&#xff0c;聚焦食堂消费、手机充值、网费缴纳等典型校园一卡通场景&#xff0c;助力理解桌面应用开发全流程。资源包共30个文件&#xff0c;含7个核心Java源码文…

作者头像 李华
网站建设 2026/9/12 14:39:30

基于HNR-gram的轴承故障诊断MATLAB实现

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 14:39:13

CDLF多级泵扬程不足的诊断与解决方案

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华