MAX Python SDK 中 log_probabilities 模块解析:ragged 批量 Token 对数概率的计算图构建与执行
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
导读
本篇文章围绕 MAX Python SDK 中max.pipelines.lib.log_probabilities模块展开,该模块负责为「批量、长度不等(ragged)」的输入序列构建并执行对数概率(log probabilities)计算图,是文本生成管线返回logprobs的核心支撑。读完本文,你将掌握log_probabilities_ragged_graph与compute_log_probabilities_ragged两个 API 的输入输出契约、底层 custom op 的图构建细节、top-k 对数概率的堆式实现约束,以及它们如何在LogProbabilitiesMixin中与 PipelineModel、OpenAI 兼容的 logprobs 语义无缝衔接。
模块定位:为 batched ragged 序列计算对数概率
max/python/max/pipelines/lib/log_probabilities.py的模块 docstring 将其职责概括为:Builds computation graphs for log probabilities over batched input sequences(为批量输入序列构建对数概率计算图)。它对外仅暴露两个公开函数:
log_probabilities_ragged_graph(device, *, levels):构建一个编译期固定的计算图;compute_log_probabilities_ragged(device, model, ...):给定已编译的模型与运行时缓冲区,真正执行计算并返回每个批次的LogProbabilities。
文档索引 pipelines.lib.log_probabilities.rst 将其列于max.pipelines.lib子模块体系之下,与arch_lookup、interfaces并列(见 pipelines.lib.rst)。这里的「ragged」指一批请求的序列长度各不相同,需要借助行偏移(row offsets)描述每个 batch 项对应的 token 区间,而不是用定长矩阵填充。
两个核心函数的 API 契约
log_probabilities_ragged_graph:一次性构建可复用的计算图
函数签名与语义来自 log_probabilities.py:
def log_probabilities_ragged_graph(device: DeviceRef, *, levels: int) -> Graphdevice:该图将要运行的设备类型;levels:期望支持的最大 top-k 的log2(max_k + 1)。例如要支持 OpenAI API 的logprobs=5,需要levels=3;更高的 levels 可支持更大的 k。
图内部按levels决定每个输出位置保留的候选列数:
out_per_token = 2**levels if levels > 0 else 1所有输入张量使用固定 dtype:logits 为DType.float32,token 与各类偏移均为DType.uint32。图中定义了 7 个输入:
| 输入 | 形状 | 含义 |
|---|---|---|
| logits | ("bseq_or_b", "vocab") | 全部 token 的 logits(echo 时为整段,否则退化为 next_token_logits) |
| tokens | ("batch_seq",) | 扁平化的 token 数组 |
| sampled_tokens | ("batch",) | 每批实际采样出的 token |
| logit_row_offsets | ("batchp1",) | 每批 logit 行的起始偏移 |
| token_row_offsets | ("batchp1",) | 每批 token 行的起始偏移 |
| lp_output_offsets | ("batchp1",) | 每批对数概率输出的行偏移(设备侧) |
| lp_output_offsets | ("batchp1",) | 同上(host 侧,供主机端索引) |
其中lp_output_offsets同时传入设备侧与 host 侧两份,是因为「输出行数由 echo 决定、且输出索引在主机端切片」,需要两端同时可见。
图的输出通过ops.custom("compute_log_probabilities_ragged", ...)指定(对应max.graph.ops的自定义算子机制):
lp_logits:形状("out_batch_seq", out_per_token),float32;lp_tokens:形状("out_batch_seq2", out_per_token),uint32。
注意源码中留有一处 TODO(GEX-2198):out_batch_seq2本应与out_batch_seq相同,但如此会让 KGEN 阶段失败,因此两个维度被拆开声明。这属于图构建层面的实现约束。
compute_log_probabilities_ragged:执行计算并组装结果
函数签名见 log_probabilities.py:
def compute_log_probabilities_ragged( device: Device, model: Model, *, input_row_offsets: npt.NDArray[np.integer[Any]], logits: Buffer | None, next_token_logits: Buffer, tokens: npt.NDArray[np.integer[Any]], sampled_tokens: npt.NDArray[np.integer[Any]], batch_top_n: Sequence[int], batch_echo: Sequence[bool], ) -> list[LogProbabilities | None]关键参数语义:
device:大部分对数概率计算所在设备,无论该参数如何设置,主机端仍会执行一小部分计算;model:必须是log_probabilities_ragged_graph构建并编译出的模型;input_row_offsets:按 batch 索引划分 token 区间的偏移数组,长度比 batch 数多 1(batch n 对应 token 索引区间[input_row_offsets[n], input_row_offsets[n+1]));logits:形状(tokens, vocab)的全量 logits,只有所有batch_echo均为 False 时才允许省略;next_token_logits:形状(batch, vocab)的下一 token logits;tokens/sampled_tokens:扁平 token 数组与每批采样 token;batch_top_n:每批要返回的 top 对数概率个数,top_n == 0的项直接跳过(返回None);batch_echo:是否在返回的对数概率中包含输入(prompt)token。
函数开头包含一整套形状与 dtype 断言(log_probabilities.py):input_row_offsets必须为一维、logits 为二维、各 batch 维度参数长度必须一致、logits 与 next_token_logits 的 vocab 维度一致,且设备侧 Buffer 必须位于指定 device 上、dtype 必须为 float32。这些断言把错误提前到「调用边界」暴露,而不是等设备端执行时才失败。
ragged 数据编排:从输入缓冲到 kernel 调用
当logits is None(即完全不 echo)时,函数走一条简化路径(log_probabilities.py):
kernel_logits = next_token_logits logit_row_offsets = np.arange(batch_size + 1, dtype=np.uint32)即直接把每批一行 next_token_logits 作为 kernel 输入,行偏移退化为0..batch_size的等差数列;否则kernel_logits = logits、logit_row_offsets = input_row_offsets。
输出行数由 echo 决定——echo 的批次输出该批所有输入 token 的对数概率,否则只输出 1 行:
output_counts = np.array([ input_row_offsets[i + 1] - input_row_offsets[i] if echo else 1 for i, echo in enumerate(batch_echo) ], dtype=np.uint32) output_row_offsets = np.concatenate( [np.zeros(1, dtype=output_counts.dtype), np.cumsum(output_counts)] )随后通过model.execute(...)一次性提交 7 个输入(token / sampled_tokens / 各类 offsets 均由 numpy 转成 uint32 Buffer 并.to(device)),得到lp_logits与lp_tokens两个输出 Buffer 并回拷到 host(log_probabilities.py)。
top-k 语义:堆式 kernel 与「采样 token 兜底」
图支持的最大 top-n
模块顶部定义了两个模块级常量:
_LOGPROBS_HEAP_LEVELS = 3 # 图构建时使用的堆深度 _MAX_TOP_LOGPROBS = 2**3 - 1 # 图可返回的最大 top-k = 7文档注释解释了原因(log_probabilities.py):kernel 内部用一个容量为2**levels - 1的最小堆维护候选,同时图会额外预留一个输出槽位给采样 token,因此该图无法支持更大的top_n。请求路由(request routes)会据此在边界校验超范围值,避免在模型 worker 内部抛错导致服务进程一起挂掉。
compute_top 的结果组装
每个输出行的 top-k 计算在主机端完成(log_probabilities.py):
if top_n < 0: raise ValueError(...) if top_n > lp_logits.shape[1] - 1: raise ValueError("top_n exceeds ... raise _LOGPROBS_HEAP_LEVELS and rerun")- 先以
token < vocab_size过滤掉填充列,将(token, logit)配对按 logit 降序排序,截断到top_n; - 特殊兜底:如果采样 token 不在 top-n 中,仍然会把它强行放入结果——这是 OpenAI 兼容 logprobs 语义的一部分(返回中必须包含实际采样 token 的对数概率)。实现上取输出行最后一列
lp_tokens[output_index, -1]与lp_logits[output_index, -1]直接写入字典,覆盖可能重复的键。
最终每个 batch 项返回一个LogProbabilities对象(top_n == 0的项返回None),其中token_log_probabilities取各行最后一列(即采样 token 的对数概率),top_log_probabilities为每行的compute_top结果列表。
数据结构:可序列化的 LogProbabilities
计算结果的承载类型定义在 max/python/max/pipelines/context/log_probabilities.py,是一个基于msgspec.Struct的纯数据类(tag=True, omit_defaults=True,便于序列化与传输):
class LogProbabilities(msgspec.Struct, tag=True, omit_defaults=True): token_log_probabilities: list[float] # 每个 token 的概率 top_log_probabilities: list[dict[int, float]] # top token 及其概率它只负责存储与序列化,不提供任何计算逻辑;该类型在 max/python/max/pipelines/context/init.py 中被 re-export,供管线各层引用。
与 PipelineModel 的集成:LogProbabilitiesMixin
max.pipelines.lib.log_probabilities还导出一个LogProbabilitiesMixin(log_probabilities.py),它要求宿主类必须是PipelineModel,且其ModelInputs子类具备tokens与input_row_offsets两个 Buffer 字段。
- 构造阶段:
__init__中取self.devices[0]作为对数概率设备,用levels=_LOGPROBS_HEAP_LEVELS构建图并通过session.load(graph)编译缓存到self._logprobs_model——因此每个模型实例只编译一次,后续每步 decode 复用; compute_log_probabilities方法:从model_outputs.next_token_logits与model_inputs中取回 numpy 数据,依据self.return_logits判断是否有全量 logits,然后委托compute_log_probabilities_ragged。
echo 的前置条件:ReturnLogits 枚举
echo 输入 token 的对数概率需要全量 logits 可用。LogProbabilitiesMixin.compute_log_probabilities中有明确的守卫(log_probabilities.py):
has_full_logits = self.return_logits in (ReturnLogits.ALL, ReturnLogits.VARIABLE) if any(batch_echo) and not has_full_logits: raise ValueError( "Log probabilities with echo=true requires enable_echo=true " "in the pipeline configuration to return logits for all tokens." )ReturnLogits是定义在 max/python/max/nn/transformer/transformer.py 的字符串枚举:LAST_TOKEN/VARIABLE/ALL。在 TextGenerationPipeline 构造函数中,模型实例化时按配置选择返回模式:
return_logits=ReturnLogits.ALL if self._pipeline_config.model.enable_echo else ReturnLogits.LAST_TOKEN也就是说:使用 echo 式对数概率(echo=true)必须先开启enable_echo=true管线配置;未开启时仅能返回LAST_TOKEN模式下的 next-token 对数概率。
在生成管线中的调用时机与数据流转
TextGenerationPipeline.execute在完成采样、拿到new_tokens之后、写回 context 之前调用对数概率计算(text_generation.py):
if inputs.enable_log_probs: with Tracer("compute_log_probabilities"): try: batch_log_probabilities.append( self._pipeline_model.compute_log_probabilities( self.session, curr_step_inputs, model_outputs, new_tokens, inputs.batch_top_log_probs, inputs.batch_echo, ) ) except NotImplementedError: logger.warning(...) batch_log_probabilities.append([None for _ in flat_batch])- 由
inputs.enable_log_probs总开关控制; NotImplementedError会被捕获并降级为整批None(不支持的模型不会中断服务);- 结果经 pipeline_variants/utils.py 的
update_context_and_prepare_responses按 batch 索引写入各 context,最终进入TextGenerationOutput; - context 侧通过
advance_token_buffer/realize_future_token将LogProbabilities存进_log_probabilities_data(见 context.py),供输出时按 token 位置取回。
值得注意的是,compute_log_probabilities的调用位于采样之后,new_tokens(即 sampled_tokens)作为参数传入,这正是「返回的 top-k 中必须包含实际采样 token」这一兜底逻辑能成立的原因。
边界条件与使用建议
- top_n 上限:默认图(
levels=3)最多支持top_n=7,覆盖 OpenAI API 的logprobs=5需求绰绰有余;需要更大 k 时必须提高_LOGPROBS_HEAP_LEVELS并重新编译,否则会在compute_top中抛ValueError。 - 与 echo 的组合:
echo=true依赖enable_echo=true配置以获得全量 logits(ReturnLogits.ALL/VARIABLE);否则必须在请求中关闭 echo。 - 类型与形状前置校验:所有断言集中在
compute_log_probabilities_ragged入口,调用方应保证 token/偏移数组使用无符号整数语义、设备端 Buffer 位于device上且为 float32,避免把错误留到 kernel 执行阶段。 - 设备与 host 分工:大头计算在
device上完成,但行偏移拼接、top 排序、采样 token 兜底等小段逻辑固定运行在 host,属设计使然而非性能缺陷。
总结
max.pipelines.lib.log_probabilities是 MAX 文本生成管线中「logprobs 能力」的最小完整实现单元:log_probabilities_ragged_graph负责把 ragged 批量对数概率计算固化为一个可复用的计算图(custom opcompute_log_probabilities_ragged),compute_log_probabilities_ragged负责执行与主机端结果组装,LogProbabilitiesMixin则把二者缝合进任意PipelineModel。理解它的输入偏移协议、levels与 top-k 的指数关系,以及 echo 与ReturnLogits的联动约束,是正确使用或扩展 MAX 管线 logprobs 功能的关键。
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考