尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

Mojo 的 max.experimental.compilation 模块指南:使用 stage / compile / as_subgraph 追踪并编译张量函数 Mojo 的 max.experimental.compilation 模块指南使用 stage / compile / as_subgraph 追踪并编译张量函数【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo本指南系统讲解 Modular PlatformMAX MojoPython 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分类符号作用Transformscompile追踪并编译先传规格得到CompiledCallable再对真实张量调用Transformsstage只追踪不编译返回StagedGraph用于检查计算图打印为 MLIRTransformsas_subgraph把被修饰/被调用的函数降级为共享子图体shared subgraph body可作装饰器或调用点使用ResultsCompiledCallable已编译的张量函数可像原函数一样被调用支持execute_raw与export_mefResultsStagedGraph追踪得到的图对象持有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, gain3.0) out run(Tensor.ones([4, 2], deviceCPU()), gain3.0) # batch accepts 4关键点有三处规格只描述张量参数。step有两个参数x是张量用TensorLayout声明gain是浮点数不进入规格在追踪时被烘焙bake进图里所以两次调用都必须传gain3.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, *, weightsNone, nameNone, custom_extensions(), allow_subgraphsTrue, signal_devices(), is_device_graphFalse)的返回是一个「接收规格」的 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(weightsweights)即「先追踪、再编译」两个阶段被封装成了一个入口。stage 与 StagedGraph编译前检查 MLIRstage(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(weightsNone)把图编译成CompiledCallable权重映射可选传入。str()/repr()输出整个模块的 MLIR。stage的容器示例来自源码 docstringdef 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, *, nameNone, prefix, key_INFER_KEY)可作装饰器也可在调用点使用。其核心价值是去重同一个函数体只被定义一次多处调用只发出mo.call。装饰器用法来自源码 docstringcompilation.as_subgraph def block(x: Tensor) - Tensor: return x * 2 spec TensorType(DType.float32, [4], deviceDeviceRef.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_placeholderTrue) def model(x: Tensor) - Tensor: for layer in (layers.0., layers.1.): x compilation.as_subgraph(block, prefixlayer)(x) return x one Tensor.ones([1], deviceCPU()) weights {layers.0.w: one * 2, layers.1.w: one * 10} run compilation.compile(model, weightsweights)(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 []), prefixprefix)把前缀在调用点传入。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 导出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], deviceCPU())) # [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类。它保存每个张量参数的一份 layoutin_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, ...)weightsNoneNone外部常量数据按图命名键控分布权重组每分片一项nameNone函数名图名称自动清洗非法字符custom_extensions()空自定义 Mojo 内核库路径allow_subgraphsTrueTrue子图共享而非内联signal_devices()空规格之外的集合通信设备is_device_graphFalseFalse是否记录设备图stage(fn, ...)同compile无weights—只追踪返回StagedGraphas_subgraph(fn, ...)nameNone函数名子图体名称prefix空串权重名前缀按调用点区分共享体权重key_INFER_KEY自动推导None时改以 IR 哈希去重StagedGraph.compile()weightsNoneNone编译为CompiledCallableCompiledCallable.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),仅供参考
返回列表