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:它提供compile、stage、as_subgraph三个转换入口,以及CompiledCallable、StagedGraph两个结果类型,用于把「输入是张量的 Python 函数」追踪(trace)成计算图并编译为可执行产物。读完本文,你将掌握如何用规格(spec)声明张量边界、两步式编译真实调用、在 MLIR 层面检查图结构、共享子图去重以及将编译产物导出为 MEF 文件。
该模块的 API 清单定义在 experimental.compilation.rst,完整实现位于 compilation.py,本文以源码中的 docstring、类型标注与测试断言为准展开。
模块总览:两个转换组、五个公开符号
max.experimental.compilation提供三个「转换函数」(Transforms)与两个「结果类型」(Results):
| 分类 | 符号 | 作用 |
|---|---|---|
| Transforms | compile | 追踪并编译:先传规格,得到CompiledCallable,再对真实张量调用 |
| Transforms | stage | 只追踪不编译,返回StagedGraph,用于检查计算图(打印为 MLIR) |
| Transforms | as_subgraph | 把被修饰/被调用的函数降级为共享子图体(shared subgraph body),可作装饰器或调用点使用 |
| Results | CompiledCallable | 已编译的张量函数,可像原函数一样被调用,支持execute_raw与export_mef |
| Results | StagedGraph | 追踪得到的图对象,持有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关键点有三处:
- 规格只描述张量参数。
step有两个参数:x是张量,用TensorLayout声明;gain是浮点数,不进入规格,在追踪时被烘焙(bake)进图里,所以两次调用都必须传gain=3.0。 - 符号维度(SymbolicDim)。
["batch", 2]中的"batch"是一个符号维度,编译后该维度接受任意大小——示例中规格声明宽度为 2,实际调用传入[4, 2]也合法。 - 类型与设备必须匹配。源码
_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.layout。compile与stage的 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.w与layers.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时,改为比较被追踪的 IR(share_subgraph对 IR 做哈希)——适合函数体内容会随环境变化的情况。 - 传入字符串时,与参数结构、操作数类型、是否有前缀拼接成完整键(
f"{key}|{name}|{structure}|{types}|{bool(prefix)}"),命中缓存则复用已生成的图体。
as_subgraph有明确的错误语义:在图捕获(capture)之外调用会抛TypeError,提示需要compile()/stage()/F.lazy()环境,此时应直接调用原函数以急切(eager)执行。
CompiledCallable:真实张量上的调用与 MEF 导出
CompiledCallable是compile的最终产物,行为与原始 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惰性初始化并缓存模型实例)。
底层机制:签名、真实化上下文与信号缓冲
从源码结构可以梳理出三条内部机制,理解它们有助于排查「为什么编译失败」:
_Signature负责边界映射(compilation.py 中的_Signature类)。它保存每个张量参数的一份 layout(in_specs)和返回值结构(out_structure)。flatten先校验参数个数与关键字集合,再用tree.paths比对 pytree 结构,然后逐个把Tensor拆成local_shards对应的 driver buffer——分布式参数按分片数展开为多个图输入。unflatten是它的逆向:把展平的图结果按out_structure重新拼成返回值。GraphRealizationContext贯穿追踪过程(见 realization_context.py 中的subgraph_context/open_subgraph/share_subgraph)。stage在realization_context(ctx), ctx:的上下文中调用被追踪函数:先用_argument_tensor把图输入重建成「未实现(unrealized)」的Tensor传给函数,函数返回后再graph.output(*flat)收尾。- 多设备集合通信需要信号缓冲(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=None | None | 外部常量数据,按图命名键控,分布权重组每分片一项 |
name=None | 函数名 | 图名称,自动清洗非法字符 | |
custom_extensions=() | 空 | 自定义 Mojo 内核库路径 | |
allow_subgraphs=True | True | 子图共享而非内联 | |
signal_devices=() | 空 | 规格之外的集合通信设备 | |
is_device_graph=False | False | 是否记录设备图 | |
stage(fn, ...) | 同compile(无weights) | — | 只追踪,返回StagedGraph |
as_subgraph(fn, ...) | name=None | 函数名 | 子图体名称 |
prefix="" | 空串 | 权重名前缀,按调用点区分共享体权重 | |
key=_INFER_KEY | 自动推导 | None时改以 IR 哈希去重 | |
StagedGraph.compile() | weights=None | None | 编译为CompiledCallable |
CompiledCallable.export_mef() | path | 必填 | 导出 MEF,无需权重与设备内存 |
TensorLayout | dtype, 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),仅供参考