资讯详情

资讯详情

MAX Graph 计算图 API 全解析:用 max.graph 在 Python 中构建、编译与运行推理图

MAX Graph 计算图 API 全解析用 max.graph 在 Python 中构建、编译与运行推理图【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojomax.graph是 Modular PlatformMAX Mojo中面向 Python 的推理图构建 API它以数据流图而非命令式逐条执行的方式描述模型的输入到输出变换让编译器能够在编译期完成形状推断、算子融合与设备编排。本文以仓库内文档 max/python/docs/graph.rst 的 API 骨架为主线结合max/python/max/graph/下的完整实现源码系统讲解 Graph / Module / KernelLibrary 三大图构建原语、Value 与 Type 体系、Shape/Dim 维度系统、设备抽象、权重与分片策略、调试配置以及子图与性能剖析等能力帮助你直接上手用max.graph编写可编译、可分发、可调试的推理图。max.graph 模块概览max.graph是 MAX 中构建推理图的 Python 包其模块 docstring 定位为 APIs to build inference graphs for MAX见 max/python/max/graph/init.py。与命令式执行不同Graph捕获操作之间的数据流MAX 在编译期据此优化并并行化执行操作运行在编译后的对象上。这种先构图、后编译、再执行的模型是 MAX 性能与设备分发的基础。根据graph.rst中的 toctree模块包含三个子模块子模块职责graph.ops全部图操作原语matmul、elementwise、reshape、reduction、control flow 等对应源码目录max/python/max/graph/ops/graph.quantization量化编码QuantizationEncoding与量化相关类型源码见max/python/max/graph/quantization.pygraph.weights权重加载含 GGUF、Safetensors 读取器源码见max/python/max/graph/weights/目录顶层命名空间在__init__.py中统一导出图构建类Graph、GraphBlock、KernelLibrary、Module、自定义扩展default_custom_extensions、default_custom_extensions_scope、图值类型Value、TensorValue、BufferValue及 Like 别名、类型系统Type、TensorType、BufferType、维度与形状Dim、StaticDim、SymbolicDim、AlgebraicDim、Shape、设备DeviceKind、DeviceRef、DevicePlacementPolicy、权重Weight、ShardingStrategy以及配置GraphDebugConfig、ProfileScopeColor。图构建核心Graph、Module 与 KernelLibraryGraph两种等价的构图方式Graph定义于 max/python/max/graph/graph.py代表单个 MAX 图。从源码 docstring 看它有两种推荐的构造策略。方式一上下文管理器。在with Graph(...) as graph:块内从graph.inputs取输入、调用ops.*构建节点、最后用graph.output(...)声明输出。块内调用的 op 会自动找到当前激活的图块外调用则会失败没有激活的图。完整的可运行示例源码中的 linear_relu 例子from max.dtype import DType from max.graph import DeviceRef, Graph, TensorType, Weight W Weight(W, DType.float32, [3, 2], DeviceRef.CPU()) b Weight(b, DType.float32, [2], DeviceRef.CPU()) with Graph( linear_relu, input_types[TensorType(DType.float32, [batch, 3], deviceDeviceRef.CPU())], ) as graph: x graph.inputs[0].tensor y x W b graph.output(y)注意这里input_types中batch是符号维度SymbolicDim3是静态维度TensorValue支持矩阵乘、等运算符重载直接拼出计算表达式。方式二构造函数传入 forward 可调用对象。将前向函数作为forward参数传给Graph构造器会自动把输入TensorValue传给该函数并把返回值记录为图输出from max.dtype import DType from max.graph import DeviceRef, Graph, TensorType, TensorValue, Weight, ops class Linear: def __init__(self, in_dim: int, out_dim: int): self.weight Weight(W, DType.float32, [in_dim, out_dim], DeviceRef.CPU()) self.bias Weight(b, DType.float32, [out_dim], DeviceRef.CPU()) def __call__(self, x: TensorValue) - TensorValue: return ops.matmul(x, self.weight) self.bias linear_layer Linear(2, 2) graph Graph( linear, linear_layer, input_types[TensorType(DType.float32, (2,), DeviceRef.CPU())], )从源码graph.py的__init__看forward分支底层依然会打开并关闭图上下文with self: result forward(*self.inputs, *args, **kwargs)然后根据返回值是None无输出、单个值还是可迭代对象多输出统一调用self.output(*outputs)。Graph构造参数汇总源码__init__签名参数默认值说明name必填图的名字forwardNone前向可调用对象传入时立即构图input_types()图输入的类型序列通常是TensorType也允许BufferType可变原地输入pathNone已保存图的路径内部使用custom_extensions[]自定义算子库路径.mojoc文件或含自定义算子的.mojo源码目录kernel_libraryNone预构建的KernelLibrary默认按custom_extensions新建moduleNone已有的 MLIR 模块内部使用传入后多个图可共享同一模块strict_device_placementDevicePlacementPolicy.Warn隐式 CPU 转移时的处理策略见后文 Devices 一节is_device_graphFalse是否标注为 device graph构图期间构造器会扫描input_types中的符号维度并登记为图的参数self._params符号维度名会写入 MLIR 的input_parameters属性保证 IR 生成的确定性。图上下文与当前图访问Graph.current类属性返回执行上下文中当前激活的图。它在with graph:进入时通过模块级ContextVar CURRENT_GRAPH设置退出时恢复。没有激活图时抛出LookupError(No graph found)。graph.inputs返回对应input_types的输入Value序列内部自动排除用于副作用排序的 chain 值。graph.output(*outputs)终结图构建并写入 MLIR 的 signature、argument_names、result_names等元数据随后对整个图做一次完整 verify。图在调用output()之前不可执行output_types属性只有在output()之后才能读取否则抛TypeError。graph.copy()深拷贝整个图不共享 MLIR 状态但共享 KernelLibrary适合把图交给其他线程做后台编译。graph.add_weight(weight, force_initial_weight_on_hostTrue)向图注册权重同名权重已存在且为同一Weight实例时返回既有值同名不同实例则抛ValueError。force_initial_weight_on_hostTrue时权重先分配在 host再转移到目标设备。Module多图共享的编译单元Module是一个或多个 Graph 的容器。它持有 MLIR module可能包含多个mo.graphop例如视觉编码器 语言模型可以作为一个编译单元交给max.engine.InferenceSession.load_all一次性编译。用同一个module参数构造多个Graph每个图都会作为该 module 的顶层 op 加入。源码中的完整示例graph.py中Moduledocstringimport numpy as np from max.driver import Accelerator, CPU, accelerator_count from max.dtype import DType from max.engine import InferenceSession from max.graph import DeviceRef, Graph, Module, ops device Accelerator() if accelerator_count() 0 else CPU() device_ref DeviceRef.from_device(device) module Module() with Graph(encoder, input_types[], modulemodule) as encoder: encoder.output( ops.constant( np.array([1.0, 2.0, 3.0], dtypenp.float32), DType.float32, devicedevice_ref, ) ) with Graph(decoder, input_types[], modulemodule) as decoder: decoder.output( ops.constant( np.array([4.0, 5.0, 6.0], dtypenp.float32), DType.float32, devicedevice_ref, ) ) models InferenceSession(devices[device]).load_all(module) (enc_out,) models[encoder.name].execute() (dec_out,) models[decoder.name].execute()关键成员mlir_module暴露底层 MLIR modulebuiltin.ModuleOp供图编译器内部与序列化工具使用常规用户无需触碰。top_level_graph_names()返回 module 中所有顶层非子图图的名字顺序与 MEF 模型顺序一致也就是InferenceSession.load_all返回模型的顺序。KernelLibrary自定义算子的加载与校验KernelLibrary管理图的内核库即自定义算子与内核的集合。它支持从两类来源加载对应load_paths的智能加载逻辑见graph.pyMojo 预编译二进制包.mojoc扩展直接add_path加入分析。Mojo 源码目录先调用_build_mojo_source_package把源码目录构建成预编译 Mojo 文件再加入分析。既不是二进制包也不是源码目录的路径会直接抛ValueError。加载后的库路径会以_kernel_library_paths属性写入图的 MLIR op从而在序列化/反序列化时保留。常用接口library_paths()返回当前已加载的内核库路径列表。add_path(path)向分析添加一个库路径。load_paths(custom_extensions)从多个路径批量加载自定义算子。__getitem__(kernel)/__contains__(kernel)/__iter__()按名字查内核、判断存在性、按序枚举符号名。verify_custom_op(custom_op)在 MLIR 上下文中校验自定义 op 的合法性带缓存相同路径集合复用Analysis。has_shape_function(kernel)返回某内核是否注册了 shape 函数。未注册 shape 函数的内核无法在运行时计算输出形状图编译器会拒绝为其声明数据依赖的输出维度内核不存在时抛KeyError。自定义扩展default_custom_extensions 与作用域default_custom_extensions()返回每个新图都会隐式加载的扩展路径元组。后端如果需要内核覆盖库kernel-overlay library会在此注册使得算子在未显式传custom_extensions的图路径上也能解析到 overlay —— 默认情况下它是空的()无任何副作用。Graph.__init__中会把这些默认路径追加在显式扩展之后去重保证即使不传custom_extensions也能触达后端的内核 overlay。default_custom_extensions_scope(*paths)是一个上下文管理器在块执行期间向默认扩展追加路径已注册的不会重复退出时恢复之前的默认值。典型用途是在一段构图代码内临时挂载自定义算子库。图值体系Value、TensorValue 与 BufferValueValue是图内所有符号值的基类见 max/python/max/graph/value.py可代表某个节点的输出、图的输入或图内任意可用的符号值。概念上可以把Value看作数据流图中的一条边。它是一个抽象类不能直接构造需要通过Value.from_mlir按 MLIR 类型分派到具体子类TensorValuemo.TensorType值语义张量 —— 每次操作产生新值BufferValuemo.BufferType可变语义的原地内存引用_ChainValuemo.ChainType副作用排序链用户通常无需直接接触_OpaqueValuemo.OpaqueType不透明值自定义算子。Value基类提供tensor/buffer/opaque三个属性做类型窄化类型不对时抛TypeError以及to_mlir()转换。TensorValue图内值语义张量TensorValue是构图时打交道最多的类型源码中为它实现了完整的张量运算面形状操作reshape(shape)、flatten(start_dim0, end_dim-1)、broadcast_to(shape)、permute(dims)、transpose(dim_1, dim_2)、T等价transpose(-1, -2)、rebind(shape, message)。其中rebind在运行时断言张量形状与给定shape一致可用于约束动态维度、提供形状提示rank 不匹配时抛ValueError失败时携带可选message报错。类型与归约cast(dtype)、argmax(axis-1)、max(axis-1)、mean(axis-1)、min(axis-1)、stdev(axis-1)、var(axis-1)布尔张量上var抛TypeError因为方差对布尔值无定义。设备与 I/Oto(device)在图中插入图执行期的设备转移节点等价ops.transfer_to用于在forward内部路由激活张量它区别于max.experimental.nn.Module.to(device)那种编译前的 host 侧权重搬移。另有print(labeldebug_tensor)打印调试信息。切片与运算符重载__getitem__支持整数、slice、ellipsis 与索引元组转成ops.slice_tensor算术与逻辑运算符全套重载 - * / // % ** | ^ ~及 ! 全部映射到ops下的对应原语__contains__/__iter__被显式禁用避免误落到__getitem__造成死循环。BufferValue可变状态与 KV CacheBufferValue是可原地更新的张量内存引用这是图在多次执行之间携带状态例如 KV cache的方式。典型使用模式是ops.buffer_load读入值语义张量、计算、再ops.buffer_store写回同一块内存。它还实现了__getitem__/__setitem__便捷切片读写以及print(labeldebug_buffer)。源码 docstring 中的完整示例from max.dtype import DType from max.graph import BufferType, DeviceRef, Graph, ops buffer_type BufferType(DType.float32, shape[4], deviceDeviceRef.CPU()) with Graph(buffer_demo, input_types[buffer_type]) as graph: state graph.inputs[0].buffer # 读入值语义张量 current ops.buffer_load(state) # 把更新后的张量写回同一块内存 ops.buffer_store(state, current 1) graph.output(current)协议与类型别名HasTensorValue/HasBufferValue是runtime_checkable协议任何实现__tensorvalue__()/__buffervalue__()的对象都可以隐式转为TensorValue/BufferValueWeight即通过该协议参与运算。TensorValueLike TensorValue | Shape | Dim | HasTensorValue | NumericNumeric 含 Python/numpy 标量及DLPackArrayBufferValueLike BufferValue | HasBufferValue。这两个 Like 别名正是文档中列出的TensorValueLike、BufferValueLike广泛用于 op 的参数类型标注。类型系统Type、TensorType 与 BufferType图内每个值都有类型由Type表示见 max/python/max/graph/type.py。Type是抽象基类Type.from_mlir按 MLIR 类型分派到TensorType、BufferType、_ChainType、_OpaqueType。TensorType符号化张量类型TensorType(dtype, shape, device)纯符号描述张量在计算中某一点的元素类型、形状与目标设备不持有任何数据。编译器在构图期用它做形状推断与优化实例可直接传给InferenceSession.load或Module.compile实验性定义图/模型的输入类型。from max.graph import TensorType, DeviceRef from max.dtype import DType tensor_type TensorType(DType.float32, (2, 3), deviceDeviceRef.CPU()) print(tensor_type.dtype) # DType.float32 print(tensor_type.shape) # [2, 3]要点shape 支持静态整数、符号字符串、代数符号表达式三种维度rank 在构图期即已知静态维度必须非负源码__init__校验否则抛TypeError。rank、num_elements()含符号维度时抛RuntimeError、cast(dtype)同形状换 dtype、parameters所依赖的符号维度名、as_buffer()转成对应 BufferType。_layout字段携带卷积滤波器的内存布局FilterLayoutRSCF / QRSCF / FCRS / FCQRS / CFRS以layout元数据写入 MLIR 类型。BufferType可变张量引用BufferType是可原地修改的张量引用的符号类型构造参数与TensorType相同dtype/shape/device继承同一基类_TensorTypeBase的rank、num_elements、cast等能力并提供as_tensor()反向转换。图输入中同时允许TensorType只读输入与BufferType可变 in-place 输入。不透明类型与链类型_OpaqueType(name, parameters)自定义算子的不透明类型参数值支持bool/int/str/DType四种源码_value_to_attribute逐一映射为 MLIR attribute反序列化由_attribute_to_value还原。_ChainType用于对副作用操作排序buffer_load、buffer_store、buffer_store_slice等都消费当前 chain 并产出新 chain每条 chain 至多使用一次从而避免数据竞争。作为用户通常无需直接构造它。维度与形状Dim、Shape 及三种维度子类Dim见 max/python/max/graph/dim.py描述张量的一个维度。通常不需要直接构造Dim直接把维度值传给TensorType/BufferType即可Dim.__new__会自动按输入类型分派输入维度类型说明int/np.integer/IntegerAttr/SIMDAttrStaticDim编译期即可确定的大小允许最激进的优化str/ParamDeclRefAttrSymbolicDim按名字标识的未知大小同名维度被视为相等编译器可据此优化ParamOperatorAttrAlgebraicDim由符号维度经算术表达式推导出的维度DimLike int | str | Dim | np.integer | TypedAttr。StaticDim固定大小。值域限制为-2**63 dim 2**63越界抛ValueError。静态维度越多编译期形状求解与优化机会越大。SymbolicDim命名规则为^[a-zA-Z_]\w*$非法名字抛ValueError。文档特别提醒同名符号维度会被解释为同一维把两个实际上不同的维度命名为同一个名字很可能导致编译失败因此命名必须谨慎。AlgebraicDim对符号维度做算术、*、//、一元负号等自动产生等价表达式会化简为同一形式例如Dim(x) 1 1 Dim(x) 2为真。注意限制代数维度在图内部合法但不能出现在图的输入或输出类型中因为其底层值可能有歧义如foo * bar可由多组foo、bar组合满足。Shape见 max/python/max/graph/shape.py是list[Dim]的子类提供rank、static_dims所有静态维的整数值列表、parameters依赖的符号维、is_static(shape)所有维均为非负静态维时返回真以及to_mlir/from_mlir。ShapeLike Iterable[DimLike]因此构造TensorType(DType.float32, [batch, 3], ...)时字符串与整数可混用。设备抽象DeviceKind、DeviceRef 与 DevicePlacementPolicyDeviceKind是设备类型枚举仅含CPU与GPU支持from_string解析未知值抛ValueError。DeviceRef是符号化设备引用由DeviceKind加id组成是 MLIR 设备属性的直接表示from max.graph import DeviceRef gpu_device DeviceRef.GPU() # gpu:0默认 id0 cpu_device DeviceRef.CPU(id1) # cpu:1构造时id max(id, 0)负 id 被截为 0。常用成员CPU(id0)/GPU(id0)工厂方法、is_cpu()/is_gpu()、to_device()转换为具体驱动Device、from_device(device)把max.driver的Device或DeviceRef统一转为DeviceRef、from_mlir/to_mlir双向转换。DeviceRef的相等性与哈希基于(device_type, id)不可变身份。DevicePlacementPolicy控制某 op 隐式把张量转移到 CPU这一行为如何被报告某些 op 只有 CPU 内核必须先把非 CPU 张量转移过去再执行。三个取值源码graph.py中枚举定义Ignore静默转移不输出任何提示Warn默认发出UserWarning指明 op 名称Error抛ValueError把隐式转移变成硬性的构图期失败。通过Graph(..., strict_device_placementDevicePlacementPolicy.Error)传入。从Graph的实现看构图期还会维护一个_DeviceChainMapdevice_chains[DeviceRef.CPU()]是 host 编排链debug.print、mo.call、collective 的 host 侧等推进其余条目是各设备的计算链新设备条目从 host 链播种从而保证设备间的执行顺序正确。权重与分片Weight 与 ShardingStrategyWeight见 max/python/max/graph/weight.py描述模型权重构造时指定名称、dtype、shape 与设备例如Weight(W, DType.float32, [3, 2], DeviceRef.CPU())并支持量化编码关联graph.quantization的QuantizationEncoding。Weight可像张量一样参与切片源码中weight[:, start:end]、weight[start:end]并通过HasTensorValue协议融入图运算。Graph.add_weight会把权重作为编译期常量mo.constant_external注入图因此权重是编译期数据而非属于某个 profile scope 的计算 op源码中显式attach_profile_scopesFalse。ShardingStrategy与分布式图执行的分片策略相关。源码中确认实现的分片函数包括col_sharding_strategy(weight, i, num_devices)按列分片对weight.shape[1]维切分剩余列不均分时前remainder个设备多分一列row_sharding_strategy(weight, i, num_devices)按行分片对weight.shape[0]维切分replicate_sharding_strategy(weight, i, num_devices)每个设备持有完整权重副本head_aware_col_sharding_strategy(...)按注意力头分布做列分片专为 head 数不能被设备数整除的输出投影权重设计且支持 NVFP4/FP4、MXFP6 等打包格式打包列数按张量自身宽度推断。分片范围统一由_compute_shard_range计算base_size, remainder divmod(shard_dim, num_devices)前remainder个设备拿到base_size 1份其余设备拿base_size份。调试与剖析配置GraphDebugConfig 与 profile_scopeGraphDebugConfig是max.engine.DebugConfig的一个窄视图通过Graph.debug暴露。它的source_tracebacks属性之所以放在Graph.debug上是因为该选项在构图阶段InferenceSession尚未存在时就被消费开启后每个图 op 的 MLIR Location 会携带 Python 调用栈方便回溯 op 来源它可以在导入后随时翻转源码_location每次读取实时配置而非 import 时缓存也可通过MODULAR_DEBUG环境变量等影响。其余调试选项仍在InferenceSession.debug上两者共享同一份全局状态。Graph.profile_scope(name, colorNone)是性能剖析辅助作用域内创建的每个 op 都会在 MLIR Location 上携带name标签下游 NVTX 追踪会以kernel_name [scope_name]后缀呈现作用域可嵌套最外层在前例如kernel_name [draft_forward/target_forward]。colorProfileScopeColor枚举modular_purple、blue、green、orange、purple、red、white、yellow仅用于可选的 in-region range 机制max-debug.profile-scope-tracing不参与 per-kernel 追踪名。使用限制作用域记录在单个模块级 ContextVar 中不要在兄弟图的profile_scope仍激活时去构造第二个嵌套图否则标签会在两个图之间泄漏。进阶能力子图、GraphBlock 与图序列化子图add_subgraph是图层面的函数把一段重复计算定义一次、在父图中多次调用。文档给出的场景是transformer 层在模型中出现 62 次将其包装为子图后编译器只需处理一次定义源码 docstring 称可显著缩短编译时间代价是子图内分配不能与外部共享峰值内存可能略增、算子融合不能跨越子图边界吞吐可能轻微下降。对于max.nn.Module模型优先使用Module.build_subgraph自动处理权重前缀。子图通过ops.call(sub, *args)调用add_subgraph内部会自动追加_ChainType()输入用于操作排序并为非 CPU 设备按devices参数补链参数。GraphBlock允许在拥有它的 op 尚未创建之前就填充一个 MLIR block 的 region body避免创建空 region → 填充 → 校验模式需要的校验暂停与失败时手动擦除。在with块内当前图的_current_block指向该 block常规ops.*会插入其中链状态被隔离block 使用自己的device_chains局部副作用不会泄漏到外层图。构造GraphBlock(arg_types)后可通过mlir_block属性把预构建 block 交给接受 block 的 op wrapper如mo.if_并用output(*values, terminatormo.YieldOp)终结mo.while的条件块需传mo.WhileConditionOp。图序列化Graph(path...)支持从磁盘加载已保存图内部使用_load_mlir会从 MLIR 文本解析模块并从_kernel_library_paths属性恢复内核库路径。Module._to_mlir_str(source_locations...)可将模块序列化为 MLIR 汇编文本启用 source_locations 时须先在构图期开启 source-traceback 捕获。小结max.graph构成了 MAX Python 推理栈的图前端Graph提供上下文管理器与 forward 回调两种等价构图方式Module支持多图联合编译KernelLibrary打通 Mojo 自定义算子.mojoc二进制包或.mojo源码目录TensorValue/BufferValue分别覆盖值语义计算与 KV cache 等可变状态Dim三态维度系统静态 / 符号 / 代数与TensorType符号类型支撑编译期形状推断DeviceRef与DevicePlacementPolicy管理设备放置和隐式转移策略Weight与分片策略为分布式执行铺路GraphDebugConfig、profile_scope与子图机制则覆盖调试、剖析与编译优化。进一步探索建议直接阅读 max/python/max/graph/graph.py 的类与完整 docstring、max/python/max/graph/ops/ 目录下的算子实现以及 max/python/docs/graph.rst 关联的graph.ops.rst、graph.quantization.rst、graph.weights.rst三份子模块文档页。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →