Mojo 的 max.experimental.compilation 模块指南:使用 stage / compile / as_subgraph 追踪并编译张量函数
2026/9/12 16:17:30 网站建设 项目流程

Mojo 的 max.experimental.compilation 模块指南:使用 stage / compile / as_subgraph 追踪并编译张量函数

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

本指南系统讲解 Modular Platform(MAX & Mojo)Python API 中的实验性编译模块max.experimental.compilation:它提供compilestageas_subgraph三个转换入口,以及CompiledCallableStagedGraph两个结果类型,用于把「输入是张量的 Python 函数」追踪(trace)成计算图并编译为可执行产物。读完本文,你将掌握如何用规格(spec)声明张量边界、两步式编译真实调用、在 MLIR 层面检查图结构、共享子图去重以及将编译产物导出为 MEF 文件。

该模块的 API 清单定义在 experimental.compilation.rst,完整实现位于 compilation.py,本文以源码中的 docstring、类型标注与测试断言为准展开。

模块总览:两个转换组、五个公开符号

max.experimental.compilation提供三个「转换函数」(Transforms)与两个「结果类型」(Results):

分类符号作用
Transformscompile追踪并编译:先传规格,得到CompiledCallable,再对真实张量调用
Transformsstage只追踪不编译,返回StagedGraph,用于检查计算图(打印为 MLIR)
Transformsas_subgraph把被修饰/被调用的函数降级为共享子图体(shared subgraph body),可作装饰器或调用点使用
ResultsCompiledCallable已编译的张量函数,可像原函数一样被调用,支持execute_rawexport_mef
ResultsStagedGraph追踪得到的图对象,持有graph属性,str()输出整个模块的 MLIR

模块的设计意图在源码开头的模块 docstring 中写得很清楚:compile用两次调用完成「函数到已编译函数」的转换——第一次为每个张量参数传入一个规格(dtype、shape、device),第二次在真实张量上执行;非张量参数在追踪期间被固定(fix)为常量,因此两次调用必须以相同方式传入。stage则在追踪完成后停止,便于检查图结构。

核心工作流:两步式 compile

先看模块 docstring 给出的最小示例:

from max.driver import CPU from max.dtype import DType from max.experimental import compilation from max.experimental.sharding import TensorLayout from max.experimental.tensor import Tensor def step(x: Tensor, *, gain: float) -> Tensor: return x * gain x_spec = TensorLayout(DType.float32, ["batch", 2], CPU()) run = compilation.compile(step)(x_spec, gain=3.0) out = run(Tensor.ones([4, 2], device=CPU()), gain=3.0) # "batch" accepts 4

关键点有三处:

  1. 规格只描述张量参数step有两个参数:x是张量,用TensorLayout声明;gain是浮点数,不进入规格,在追踪时被烘焙(bake)进图里,所以两次调用都必须传gain=3.0
  2. 符号维度(SymbolicDim)["batch", 2]中的"batch"是一个符号维度,编译后该维度接受任意大小——示例中规格声明宽度为 2,实际调用传入[4, 2]也合法。
  3. 类型与设备必须匹配。源码_Signature.flatten会逐一核对实参的 dtype、维度数、设备 mesh,以及符号维度之外每个静态维的大小,不匹配会抛出ValueError(见 compilation.py 中flatten的校验逻辑)。

模块的 doctest 断言了这个结果:out.to_numpy()np.full((4, 2), 3.0)全等,验证了「张量乘以标量」的追踪正确性。

compile 的完整参数

compile(fn, *, weights=None, name=None, custom_extensions=(), allow_subgraphs=True, signal_devices=(), is_device_graph=False)的返回是一个「接收规格」的 callable;对它传入每个张量参数一个规格后,得到CompiledCallable。各参数含义:

  • weights:图为外部常量(external constants)声明的数据,以图内命名作为键;分布式权重每个分片(shard)一项。
  • name:图的名字,默认取fn.__name__(源码_sanitized_graph_name会把非标识符字符折叠为下划线)。
  • custom_extensions:自定义 Mojo 内核库的路径列表。
  • allow_subgraphs:是否允许as_subgraph的函数体成为共享子图(而不是内联进调用者),默认True
  • signal_devices:除规格覆盖的设备之外、参与集合通信(collective)的设备。
  • is_device_graph:是否记录设备图(device graph)。

从源码看,compile内部就是对stage(...)的返回值调用.compile(weights=weights),即「先追踪、再编译」两个阶段被封装成了一个入口。

stage 与 StagedGraph:编译前检查 MLIR

stage(fn, ...)的参数与compile一致(不含weights),但不触发编译。它返回StagedGraph,打印它即可查看整个模块的 MLIR(包含共享子图体):

def scale(x: Tensor) -> Tensor: return x * 2 spec = TensorLayout(DType.float32, [4], CPU()) staged = compilation.stage(scale)(spec) print(staged) # 输出包含 mo.mul 的 MLIR

源码中StagedGraph.__str__返回str(self.graph._module),其注释特别指出:图的 op 只会引用子图体的名字而不包含其内容,因此渲染整个模块才能保证「检查某 op 是否缺失」这类断言不被漏检。__repr__也指向__str__,避免默认 repr 让「op 不存在」的测试误判。

StagedGraph的公开字段与行为:

  • graph:记录的Graph对象。
  • compile(weights=None):把图编译成CompiledCallable,权重映射可选传入。
  • str()/repr():输出整个模块的 MLIR。

stage的容器示例(来自源码 docstring):

def combine(kv: dict[str, Tensor], alpha: float) -> Tensor: return (kv["a"] + kv["b"]) * alpha spec = TensorLayout(DType.float32, [2], CPU()) # 容器中的每个张量都是图输入;alpha 被烘焙进图 staged = compilation.stage(combine)({"a": spec, "b": spec}, 2.0) print(staged)

对应断言len(staged.graph.inputs) == 2说明:字典里两个张量分别成为独立的图输入。pytree(这里是 dict)中的每个张量都会被展平为单独输入。

用 TensorLayout 声明边界:dtype、全局形状与设备网格

规格(spec)是本模块的核心概念。TensorLayout定义在 sharding/types.py,其字段为:

  • dtype:元素数据类型(如DType.float32)。
  • shape全局形状,不是某个分片的形状;可含符号维度。
  • device:放置方式——单个设备、用于复制的网格(mesh)或DeviceMapping;若沿网格某轴切分一个形状中不存在的轴,会抛ValueError

TensorLayout的分片语义值得注意:local_types按 mesh 顺序为每个设备生成一个TensorType;一个沿 mesh 轴切分的符号维度,在每个分片上会变成新的局部维度,命名为"{original}_{axis_name}_{shard}",从而保证「同一全局维度沿不同轴切分」在图中可区分。

TensorLayout还有两个派生/转换入口:

  • BufferLayout(TensorLayout)可写边界。以它声明的参数在每一层边界都会降级为BufferValue,因此对它的写入能传回调用者;TensorLayout本身只读。可用layout.as_buffer()转换。
  • as_layout():把声明归一化为 layout。它接受TensorLayout/BufferLayout原样通过,也接受单设备TensorType/BufferType并自动强转(BufferType强转为BufferLayout以保留「可写」声明)。

边界只能由 layout 声明:源码中as_layout会拒绝活的Tensor,因为活张量的维度是它当前持有的值,直接取用会把每个维度都固定死;若确实要以某个张量的当前形状为规格,应显式传tensor.layoutcompilestage的 docstring 均强调了这一规则。

as_subgraph:共享子图体与按调用点区分权重

as_subgraph(fn, *, name=None, prefix="", key=_INFER_KEY)可作装饰器,也可在调用点使用。其核心价值是去重:同一个函数体只被定义一次,多处调用只发出mo.call

装饰器用法(来自源码 docstring):

@compilation.as_subgraph def block(x: Tensor) -> Tensor: return x * 2 spec = TensorType(DType.float32, [4], device=DeviceRef.CPU()) staged = compilation.stage(lambda x: block(block(block(x))))(spec)

对应的测试断言:str(staged).count("mo.graph @block") == 1(图体只定义一次)且str(staged).count("mo.call @block") == 3(被调用三次)。

共享体也共享其声明的权重。当同一函数体在不同调用点需要不同的权重时,用prefix为每个调用点挂出自己的权重命名空间:

w_type = TensorLayout(DType.float32, [1], CPU()) def block(x: Tensor) -> Tensor: return x * F.constant_external("w", w_type, is_placeholder=True) def model(x: Tensor) -> Tensor: for layer in ("layers.0.", "layers.1."): x = compilation.as_subgraph(block, prefix=layer)(x) return x one = Tensor.ones([1], device=CPU()) weights = {"layers.0.w": one * 2, "layers.1.w": one * 10} run = compilation.compile(model, weights=weights)(w_type) out = run(one) # [20.0]

这里prefix会被前置到函数体声明的相对权重名前,于是同一个block体在两个调用点分别解析出layers.0.wlayers.1.w,最终1 * 2 * 10 = 20.0(对应 doctest 断言)。源码中ops.call(subgraph, *operands, *(ctx.signal_buffers or []), prefix=prefix)把前缀在调用点传入。

key参数控制「什么标识同一个函数体」:

  • 省略时自动从fn推导(_inferred_key:闭包、绑定方法或带默认值的函数无法安全推导,此时返回None,改用 IR 哈希比较)。
  • 传入None时,改为比较被追踪的 IRshare_subgraph对 IR 做哈希)——适合函数体内容会随环境变化的情况。
  • 传入字符串时,与参数结构、操作数类型、是否有前缀拼接成完整键(f"{key}|{name}|{structure}|{types}|{bool(prefix)}"),命中缓存则复用已生成的图体。

as_subgraph有明确的错误语义:在图捕获(capture)之外调用会抛TypeError,提示需要compile()/stage()/F.lazy()环境,此时应直接调用原函数以急切(eager)执行。

CompiledCallable:真实张量上的调用与 MEF 导出

CompiledCallablecompile的最终产物,行为与原始 Python 函数一致——每个规格位置放一个真实张量

def scale(x: Tensor) -> Tensor: return x * 2 spec = TensorLayout(DType.float32, [3], CPU()) run = compilation.compile(scale)(spec) # 此处在编译 run.export_mef("scale.mef") # 无需权重、无需设备内存 out = run(Tensor.ones([3], device=CPU())) # [2.0, 2.0, 2.0]

其公开接口包括:

  • __call__(*args, **kwargs):以真实张量调用。内部先经_signature.flatten校验并展平参数,再execute_raw执行,最后按out_structure还原返回值;参数类型错误抛TypeError,结构与规格不符抛ValueError
  • execute_raw(*buffers) -> list[Buffer]:直接在原始缓冲上执行,返回展平的结果缓冲列表。这是「绕过规格校验、手动喂缓冲」的底层路径(flatten的错误信息也提示「use execute_raw()」)。
  • export_mef(path):把编译产物写为 MEF 文件。MEF 是运行时执行的二进制格式;导出过程不绑定权重、不分配设备内存,因此第一次调用之前就可以导出。之后可用max.engine.read读回,跳过再次编译。
  • weights:图为外部常量声明的数据映射(键为图内命名,分布式权重每个分片一项)。首次调用时绑定权重并分配设备内存(源码通过functools.cached_property _engine_model惰性初始化并缓存模型实例)。

底层机制:签名、真实化上下文与信号缓冲

从源码结构可以梳理出三条内部机制,理解它们有助于排查「为什么编译失败」:

  1. _Signature负责边界映射(compilation.py 中的_Signature类)。它保存每个张量参数的一份 layout(in_specs)和返回值结构(out_structure)。flatten先校验参数个数与关键字集合,再用tree.paths比对 pytree 结构,然后逐个把Tensor拆成local_shards对应的 driver buffer——分布式参数按分片数展开为多个图输入unflatten是它的逆向:把展平的图结果按out_structure重新拼成返回值。
  2. GraphRealizationContext贯穿追踪过程(见 realization_context.py 中的subgraph_context/open_subgraph/share_subgraph)。stagerealization_context(ctx), ctx:的上下文中调用被追踪函数:先用_argument_tensor把图输入重建成「未实现(unrealized)」的Tensor传给函数,函数返回后再graph.output(*flat)收尾。
  3. 多设备集合通信需要信号缓冲(signal buffers)_signal_device_ids汇总输入 layout 的 mesh 设备与signal_devices中参与通信的加速器;当参与设备少于两个时返回空元组。若需要,stage会把_cached_signal_buffers(ids)的类型追加进图输入,CompiledCallable则通过_cached_signal_buffers一次性分配缓冲并缓存(「allocated once, not per call」),避免每次调用重复分配。

参数速查表

符号关键参数默认值说明
compile(fn, ...)weights=NoneNone外部常量数据,按图命名键控,分布权重组每分片一项
name=None函数名图名称,自动清洗非法字符
custom_extensions=()自定义 Mojo 内核库路径
allow_subgraphs=TrueTrue子图共享而非内联
signal_devices=()规格之外的集合通信设备
is_device_graph=FalseFalse是否记录设备图
stage(fn, ...)compile(无weights只追踪,返回StagedGraph
as_subgraph(fn, ...)name=None函数名子图体名称
prefix=""空串权重名前缀,按调用点区分共享体权重
key=_INFER_KEY自动推导None时改以 IR 哈希去重
StagedGraph.compile()weights=NoneNone编译为CompiledCallable
CompiledCallable.export_mef()path必填导出 MEF,无需权重与设备内存
TensorLayoutdtype, shape, device必填全局形状 + 设备放置;只读边界
BufferLayout继承TensorLayout可写边界,写回调用者

适用范围与注意事项

  • 该模块路径带experimental,属于实验性 API,接口可能随版本演进;本文以当前仓库 compilation.py 的实现为准。
  • 非张量参数在追踪期被固定,因此compile/stage两次调用(声明规格、真实执行)中非张量参数必须保持一致;若要改变它们,需要重新编译。
  • 只能以 layout 声明边界,直接传入活Tensor会被拒绝;BufferLayout才是声明「可写」的方式。
  • 分片(sharding)场景下,实参的分片数与设备网格必须与规格一致,否则flatten会抛ValueError
  • MEF 导出不依赖权重与设备内存,适合在首次调用前完成;读回请使用max.engine.read

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询