
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.pydef log_probabilities_ragged_graph(device: DeviceRef, *, levels: int) - Graphdevice该图将要运行的设备类型levels期望支持的最大 top-k 的log2(max_k 1)。例如要支持 OpenAI API 的logprobs5需要levels3更高的 levels 可支持更大的 k。图内部按levels决定每个输出位置保留的候选列数out_per_token 2**levels if levels 0 else 1所有输入张量使用固定 dtypelogits 为DType.float32token 与各类偏移均为DType.uint32。图中定义了 7 个输入输入形状含义logits(bseq_or_b, vocab)全部 token 的 logitsecho 时为整段否则退化为 next_token_logitstokens(batch_seq,)扁平化的 token 数组sampled_tokens(batch,)每批实际采样出的 tokenlogit_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)float32lp_tokens形状(out_batch_seq2, out_per_token)uint32。注意源码中留有一处 TODOGEX-2198out_batch_seq2本应与out_batch_seq相同但如此会让 KGEN 阶段失败因此两个维度被拆开声明。这属于图构建层面的实现约束。compute_log_probabilities_ragged执行计算并组装结果函数签名见 log_probabilities.pydef 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 数多 1batch n 对应 token 索引区间[input_row_offsets[n], input_row_offsets[n1])logits形状(tokens, vocab)的全量 logits只有所有batch_echo均为 False 时才允许省略next_token_logits形状(batch, vocab)的下一 token logitstokens/sampled_tokens扁平 token 数组与每批采样 tokenbatch_top_n每批要返回的 top 对数概率个数top_n 0的项直接跳过返回Nonebatch_echo是否在返回的对数概率中包含输入prompttoken。函数开头包含一整套形状与 dtype 断言log_probabilities.pyinput_row_offsets必须为一维、logits 为二维、各 batch 维度参数长度必须一致、logits 与 next_token_logits 的 vocab 维度一致且设备侧 Buffer 必须位于指定 device 上、dtype 必须为 float32。这些断言把错误提前到「调用边界」暴露而不是等设备端执行时才失败。ragged 数据编排从输入缓冲到 kernel 调用当logits is None即完全不 echo时函数走一条简化路径log_probabilities.pykernel_logits next_token_logits logit_row_offsets np.arange(batch_size 1, dtypenp.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) ], dtypenp.uint32) output_row_offsets np.concatenate( [np.zeros(1, dtypeoutput_counts.dtype), np.cumsum(output_counts)] )随后通过model.execute(...)一次性提交 7 个输入token / sampled_tokens / 各类 offsets 均由 numpy 转成 uint32 Buffer 并.to(device)得到lp_logits与lp_tokens两个输出 Buffer 并回拷到 hostlog_probabilities.py。top-k 语义堆式 kernel 与「采样 token 兜底」图支持的最大 top-n模块顶部定义了两个模块级常量_LOGPROBS_HEAP_LEVELS 3 # 图构建时使用的堆深度 _MAX_TOP_LOGPROBS 2**3 - 1 # 图可返回的最大 top-k 7文档注释解释了原因log_probabilities.pykernel 内部用一个容量为2**levels - 1的最小堆维护候选同时图会额外预留一个输出槽位给采样 token因此该图无法支持更大的top_n。请求路由request routes会据此在边界校验超范围值避免在模型 worker 内部抛错导致服务进程一起挂掉。compute_top 的结果组装每个输出行的 top-k 计算在主机端完成log_probabilities.pyif 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的纯数据类tagTrue, omit_defaultsTrue便于序列化与传输class LogProbabilities(msgspec.Struct, tagTrue, omit_defaultsTrue): 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 的集成LogProbabilitiesMixinmax.pipelines.lib.log_probabilities还导出一个LogProbabilitiesMixinlog_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.pyhas_full_logits self.return_logits in (ReturnLogits.ALL, ReturnLogits.VARIABLE) if any(batch_echo) and not has_full_logits: raise ValueError( Log probabilities with echotrue requires enable_echotrue in the pipeline configuration to return logits for all tokens. )ReturnLogits是定义在 max/python/max/nn/transformer/transformer.py 的字符串枚举LAST_TOKEN/VARIABLE/ALL。在 TextGenerationPipeline 构造函数中模型实例化时按配置选择返回模式return_logitsReturnLogits.ALL if self._pipeline_config.model.enable_echo else ReturnLogits.LAST_TOKEN也就是说使用 echo 式对数概率echotrue必须先开启enable_echotrue管线配置未开启时仅能返回LAST_TOKEN模式下的 next-token 对数概率。在生成管线中的调用时机与数据流转TextGenerationPipeline.execute在完成采样、拿到new_tokens之后、写回 context 之前调用对数概率计算text_generation.pyif 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最终进入TextGenerationOutputcontext 侧通过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 上限默认图levels3最多支持top_n7覆盖 OpenAI API 的logprobs5需求绰绰有余需要更大 k 时必须提高_LOGPROBS_HEAP_LEVELS并重新编译否则会在compute_top中抛ValueError。与 echo 的组合echotrue依赖enable_echotrue配置以获得全量 logitsReturnLogits.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_raggedcompute_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),仅供参考