MLX 编译指南:使用 mx.compile 合并计算图、融合算子并加速训练
【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx
MLX(Array framework for Apple silicon)在mlx.core中提供了compile这一函数变换(function transformation),用于编译计算图。函数编译通过合并重复的公共计算、融合特定算子来得到更小的计算图,在多数场景下能显著降低运行时间与内存占用。本文以 compile.rst 为主体,结合 mlx/compile.cpp 与 python/src/transforms.cpp 的源码实现,系统讲解mx.compile的使用方式、缓存与重编译行为、纯函数约束、训练图编译、与其它变换的组合以及 shapeless 编译,帮助你写出可安全编译、可复用、可调试的高性能 MLX 代码。
一、编译的基本用法:从普通函数到编译函数
mx.compile的入门用法非常简单:把一个以数组为输入、数组为输出的普通 Python 函数包装起来即可。下面的例子演示了普通调用与编译调用的区别:
def fun(x, y): return mx.exp(-x) + y x = mx.array(1.0) y = mx.array(2.0) # 普通调用,不编译 # 输出: array(2.36788, dtype=float32) print(fun(x, y)) # 编译该函数 compiled_fun = mx.compile(fun) # 输出: array(2.36788, dtype=float32) print(compiled_fun(x, y))普通函数与编译函数的输出在数值精度范围内是一致的——compile是语义保持(semantics-preserving)的变换,只改变执行的组织方式,不改变计算结果。
从底层看,mx.compile返回的编译函数在 mlx/compile.cpp 中实现为一条完整的编译流水线,首次调用时会依次执行:
compile_trace:用占位符(placeholder)输入调用原函数,追踪出计算图;compile_dfs:深度优先遍历计算图,构建 tape(操作序列)与 parents map(父节点映射);compile_simplify:简化 tape——合并相同标量、移除 no-op(Copy、StopGradient)、多轮合并等价子表达式(默认 3 轮,见 mlx/compile.cpp);compile_fuse:将可融合的子图替换为Compiled原语(见 mlx/compile.cpp);compile_replace:把占位符替换为真实数组,供后续求值。
缓存机制:编译一次,多次复用
编译本身有成本:第一次调用编译函数时,MLX 需要构建计算图、优化并生成/编译代码,这个过程相对较慢。但 MLX 会对编译结果做缓存,后续调用不会重新触发编译。因此,建议只编译你打算多次调用的函数:
def fun(x, y): return mx.exp(-x) + y x = mx.array(1.0) y = mx.array(2.0) compiled_fun = mx.compile(fun) # 此处发生编译 compiled_fun(x, y) # 不再编译 compiled_fun(x, y) # 不再编译 mx.compile(fun)(x, y)缓存在 mlx/compile.cpp 的CompileCache中实现:它以函数地址(fun_id)为主键,逐条缓存 entry;命中条件包括默认 stream/设备匹配、shapeless标志一致、输入数组的 shape 与 dtype 相等以及编译常量一致。缓存使用共享锁实现线程安全,并且是线程局部的(thread_local,见 mlx/compile.cpp)。
何时会触发重新编译
以下情况会导致函数被重新编译:
- 输入的形状或维度数量发生变化;
- 任一输入的类型发生变化;
- 函数的输入个数发生变化。
其中,部分情况只重跑编译栈中的某几层(例如仅改变形状),而另一些情况(例如改变类型)会重跑整条编译栈。一般来说,应避免过于频繁地编译函数。相应地,Python 绑定层在 python/src/transforms.cpp 的compile签名中提供了三个可选参数:inputs、outputs与shapeless,其中shapeless正是为缓解形状变化导致的重编译而设计的,详见后文"Shapeless 编译"一节。
需要警惕的反模式:在循环中创建并销毁编译函数
另一个容易踩坑的写法是编译那些被频繁创建又销毁的函数,例如在循环内编译匿名函数:
a = mx.array(1.0) # 不要这样做:每一轮迭代都会重新编译这个 lambda for _ in range(5): mx.compile(lambda x: mx.exp(mx.abs(x)))(a)由于每次迭代都会新建一个 lambda 对象,其函数地址(缓存键)随之变化,缓存无法命中,从而反复编译。应把编译函数提取到循环之外,只创建一次。
二、真实加速案例:编译 GELU
mlx.nn.gelu是 Transformer 类模型中常用的非线性激活函数,其实现涉及多个一元(unary)与二元(binary)逐元素操作:
def gelu(x): return x * (1 + mx.erf(x / math.sqrt(2))) / 2- 当输入数组很小时,该函数受调用开销(overhead-bound)限制;
- 当输入数组很大时,则受内存带宽(memory bandwidth-bound)限制。
而gelu中的所有操作都是可融合的,mx.compile可以将其融合进单个 kernel,从而在这两种场景下都获得可观加速。
下面用计时辅助函数对比普通函数与编译函数的运行时间(该辅助函数先做 10 次 warm-up,并在计时循环中通过mx.eval完成同步):
import time def timeit(fun, x): # warm up for _ in range(10): mx.eval(fun(x)) tic = time.perf_counter() for _ in range(100): mx.eval(fun(x)) toc = time.perf_counter() tpi = 1e3 * (toc - tic) / 100 print(f"Time per iteration {tpi:.3f} (ms)")构造一个大数组并分别计时:
x = mx.random.uniform(shape=(32, 1000, 4096)) timeit(gelu, x) timeit(mx.compile(gelu), x)在 M1 Max 上,普通gelu约为 15.5 毫秒,编译后的gelu约为 3.1 毫秒,编译版本快约 5 倍(该数值来自文档在 M1 Max 上的实测,具体结果随设备与数组规模而异)。
从源码可以印证"哪些算子可融合"的判断逻辑:在 mlx/compile.cpp 中,is_unary覆盖Exp、Erf、Negative、Log、Sigmoid等一元算子,is_binary覆盖Add、Multiply、Divide、Subtract等二元算子,is_ternary覆盖Select,is_broadcast覆盖Broadcast,而is_fusable正是这四类的并集。compile_fuse在反向遍历 tape 时递归收集可融合的算子,并受两个常量约束:最大融合深度max_compile_depth = 11、最多输入数组数max_compile_arrays = 24(见 mlx/compile.cpp)。因此,像gelu这种由一元/二元逐元素算子串成的计算链,可以整体融合成一个Compiled原语,只需在编译阶段为该子图生成一个 kernel 即可。
三、调试编译函数:占位符、disable_compile 与 MLX_DISABLE_COMPILE
编译函数在首次被调用时,是用占位符输入进行追踪(tracing)的。这意味着在编译函数内部不能对数组求值(例如打印数组内容),否则会崩溃:
@mx.compile def fun(x): z = -x print(z) # 崩溃 return mx.exp(z) fun(mx.array(5.0))这是因为占位符数组只用于构建计算图,本身不携带数据。需要调试时,检查中间数组的内容非常有用,方法之一是全局禁用编译,使用mx.disable_compile()函数或设置环境变量MLX_DISABLE_COMPILE。例如,下面的代码即便fun是编译的,也不会崩溃:
@mx.compile def fun(x): z = -x print(z) # 正常 return mx.exp(z) mx.disable_compile() fun(mx.array(5.0))disable_compile/enable_compile在 mlx/compile.h 中声明,并绑定为mlx.core.disable_compile与mlx.core.enable_compile(见 python/src/transforms.cpp)。在 mlx/compile.cpp 中,compile_mode()首次初始化时会检查环境变量MLX_DISABLE_COMPILE:只要该变量被设置(即便设为0也会生效,属于"按存在与否启用"的变量),编译模式即为disabled;而运行时调用enable_compile()可以覆盖该环境变量。
环境变量的完整语义记录在 environment_variables.rst:布尔型变量用0关闭、非零整数开启,但MLX_DISABLE_COMPILE特殊——"只要存在即生效",所以设置0也会禁用编译;mlx.core.enable_compile可以覆盖它。此外,编译模式还有更细粒度的控制:CompileMode枚举(disabled、no_simplify、no_fuse、enabled)与set_compile_mode(见 mlx/compile.h)可供进阶诊断使用,例如只跳过简化或只跳过融合,以定位性能问题的来源。
测试方面,python/tests/test_compile.py 的test_enable_disable验证了:通过mx.export_to_dot导出计算图并统计节点数,禁用编译后节点数明显增多,重新启用后恢复为编译时的节点数——这是观察编译"合并/融合"效果最直观的手段。
四、纯函数约束:副作用、隐式输入与隐式输出
编译函数被设计为纯函数(pure),即不应该产生副作用。例如下面的代码会出问题:
state = [] @mx.compile def fun(x, y): z = x + y state.append(z) return mx.exp(z) fun(mx.array(1.0), mx.array(2.0)) # 崩溃! print(state)原因在于:首次调用fun后,state列表中保存的是一个占位符数组。占位符没有真实数据,只用于构建计算图,打印这样的数组会导致崩溃。
针对"编译函数内部更新外部容器"这一需求,文档给出了两种解决方案。
方案一:把 state 作为返回值输出
state = [] @mx.compile def fun(x, y): z = x + y state.append(z) return mx.exp(z), state _, state = fun(mx.array(1.0), mx.array(2.0)) # 输出 [array(3, dtype=float32)] print(state)方案二:用 outputs= 参数捕获隐式输出
有些场景下,显式返回更新后的 state 很不方便。因此mx.compile提供了outputs参数来捕获隐式输出:
from functools import partial state = [] # 告诉 compile 把 state 捕获为输出 @partial(mx.compile, outputs=state) def fun(x, y): z = x + y state.append(z) return mx.exp(z) fun(mx.array(1.0), mx.array(2.0)) # 输出 [array(3, dtype=float32)] print(state)这在编译包含"更新容器"逻辑的函数时特别有用——训练mlx.nn.Module参数时正是典型场景。在 python/src/transforms.cpp 的PyCompiledFun::call_impl中,outputs捕获的数组会以tree_flatten方式扁平化后追加到函数输出末尾,编译执行后通过tree_fill回写到原容器,从而把占位符替换为真实数组。
常量:参数列表之外的输入被视为常量
编译函数会把不在参数列表中的输入当作常量。例如:
state = [mx.array(1.0)] @mx.compile def fun(x): return x + state[0] # 输出 array(2, dtype=float32) print(fun(mx.array(1.0))) # 更新 state state[0] = mx.array(5.0) # 仍然输出 array(2, dtype=float32) print(fun(mx.array(1.0)))修改state后输出不变,因为首次编译时state[0]的值(常量)已被固化进编译产物。这也对应 mlx/compile.cpp 中的逻辑:compile_fuse会把"非输入、无 primitive、大小为 1 的标量"标记为constant_ids,并参与生成 kernel 名称;常量值本身通过constant_hasher哈希后拼入 kernel 名字(见 mlx/compile.cpp),因此常量变化会触发缓存键变化并重新编译。在 Python 层,非数组参数(float、int、str、None 以及 list/tuple/dict 的树形结构)也会被编码进constants向量参与缓存匹配(见 python/src/transforms.cpp)。
想让state的变化反映到输出,同样有两种办法。
办法一:把 state 作为显式输入传入
state = [mx.array(1.0)] @mx.compile def fun(x, state): return x + state[0] # 输出 array(2, dtype=float32) print(fun(mx.array(1.0), state)) # 更新 state state[0] = mx.array(5.0) # 输出 array(6, dtype=float32) print(fun(mx.array(1.0), state))办法二:用 inputs= 参数捕获隐式输入
from functools import partial state = [mx.array(1.0)] # 告诉 compile 把 state 捕获为输入 @partial(mx.compile, inputs=state) def fun(x): return x + state[0] # 输出 array(2, dtype=float32) print(fun(mx.array(1.0))) # 更新 state state[0] = mx.array(5.0) # 输出 array(6, dtype=float32) print(fun(mx.array(1.0)))inputs捕获的数组会被追加到实际输入之后参与编译,调用时先通过tree_fill将占位符填入捕获容器、执行后再用tree_replace还原(见 python/src/transforms.cpp)。inputs/outputs都支持 list 或 dict,可包含任意嵌套的 list、dict 与 array,非 array 的叶子节点会被忽略(见 python/src/transforms.cpp 的签名说明)。相关行为在 python/tests/test_compile.py 的test_compile_capture中有系统验证。
五、编译完整训练图:前向 + 反向 + 参数更新一步到位
本节用一个常见的训练设置示例,演示如何用mx.compile编译完整的前向、反向与参数更新流程:使用mlx.nn.Module定义模型、mlx.optimizers.Optimizer维护带状态(如动量)的优化器。
先看不编译的版本:
import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim # 4 个样本,每个 10 维特征 x = mx.random.uniform(shape=(4, 10)) # 0、1 标签 y = mx.array([0, 1, 0, 1]) # 简单的线性模型 model = nn.Linear(10, 1) # 带动量的 SGD optimizer = optim.SGD(learning_rate=0.1, momentum=0.8) def loss_fn(model, x, y): logits = model(x).squeeze() return nn.losses.binary_cross_entropy(logits, y) loss_and_grad_fn = nn.value_and_grad(model, loss_fn) # 执行 10 步梯度下降 for it in range(10): loss, grads = loss_and_grad_fn(model, x, y) optimizer.update(model, grads) mx.eval(model.parameters(), optimizer.state)要编译"更新"这一步,可以把整个更新过程放进一个函数,并用合适的inputs/outputs捕获状态。下面是编译后的相同示例:
import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim from functools import partial # 4 个样本,每个 10 维特征 x = mx.random.uniform(shape=(4, 10)) # 0、1 标签 y = mx.array([0, 1, 0, 1]) # 简单的线性模型 model = nn.Linear(10, 1) # 带动量的 SGD optimizer = optim.SGD(learning_rate=0.1, momentum=0.8) def loss_fn(model, x, y): logits = model(x).squeeze() return nn.losses.binary_cross_entropy(logits, y) # 将被捕获为输入和输出的状态 state = [model.state, optimizer.state] @partial(mx.compile, inputs=state, outputs=state) def step(x, y): loss_and_grad_fn = nn.value_and_grad(model, loss_fn) loss, grads = loss_and_grad_fn(model, x, y) optimizer.update(model, grads) return loss # 执行 10 步梯度下降 for it in range(10): loss = step(x, y) # 求值模型与优化器状态 mx.eval(state) print(loss)这里的关键点:
inputs=state把模型参数与优化器状态(含动量缓冲)作为隐式输入捕获,保证每次step都基于最新状态计算;outputs=state把更新后的参数与状态作为隐式输出捕获并写回,使model.state、optimizer.state在编译执行后携带真实数据;- 每次迭代后仍需
mx.eval(state)触发求值(MLX 是惰性求值框架),并打印 loss。
注意:如果使用的模块包含随机采样,例如
mlx.nn.Dropout,请务必把mx.random.state也纳入compile捕获的state中,即state = [model.state, optimizer.state, mx.random.state]。否则随机数状态不会在编译函数中被更新,导致每次调用得到相同的采样结果。这一行为在 python/tests/test_compile.py 的test_compile_rng系列测试中有专门覆盖:inputs=mx.random.state与outputs=mx.random.state必须同时捕获,随机状态才能在编译函数间正确传递。
提示:更多编译完整训练图的示例可参考 MLX 官方示例仓库(mlx-examples)。本文不展开外部链接,训练图编译的通用模式即"把 step 整体放进函数,用
inputs/outputs捕获可变状态"。
六、与其它函数变换的组合:可组合的变换体系
MLX 的函数变换是可组合的:你可以把任意函数变换应用到任意其它函数变换的输出上。编译经过变换的函数与预期一致:
grad_fn = mx.grad(mx.exp) compiled_grad_fn = mx.compile(grad_fn) # 输出: array(2.71828, dtype=float32) print(grad_fn(mx.array(1.0))) # 同样输出: array(2.71828, dtype=float32) print(compiled_grad_fn(mx.array(1.0)))一个需要注意的默认行为是:对编译函数再施加变换时,变换后的函数默认不会被编译(这是为了尽可能保留编译产物、避免重复编译)。如果需要编译变换后的函数,只需把它再传给mx.compile即可。例如test_compile_two_input_grad(python/tests/test_compile.py)验证了mx.compile(mx.grad(loss))与直接求梯度结果一致;test_vjp_vjp_compiled、test_vmap_compiled(python/tests/test_compile.py)也验证了编译函数与vjp、jvp、vmap组合的正确性。
你也可以编译那些内部调用了编译函数的函数。最佳实践是编译最外层的函数,给compile最大机会去优化整个计算图:
@mx.compile def inner(x): return mx.exp(-mx.abs(x)) def outer(x): inner(inner(x)) # 编译外层函数通常是更好的选择, # 因为即使内层函数已编译,外层编译仍可能更快 fun = mx.compile(outer)七、Shapeless 编译:一次编译,多变形状
默认情况下,编译函数的输入形状一旦改变就会重新编译。通过给mx.compile传入shapeless=True,可以只编译一次,然后在任意形状的输入上运行:
def fun(x, y): return mx.abs(x + y) compiled_fun = mx.compile(fun, shapeless=True) x = mx.array(1.0) y = mx.array(-2.0) # 首次调用触发编译 print(compiled_fun(x, y)) # 换用不同形状再次调用,不会重新编译 x = mx.array([1.0, -6.0]) y = mx.array([-2.0, 3.0]) print(compiled_fun(x, y))从源码看,shapeless的影响体现在两个层面:
- 缓存匹配:在
CompileCache::find中,shapeless模式下比较输入时跳过 shape 检查,只比较ndim与dtype(见 mlx/compile.cpp); - 输出形状推断:
compile_replace在shapeless模式下不再沿用追踪时记录的静态 shape,而是调用每个 primitive 的output_shapes(real_inputs)依据真实输入推断输出形状(见 mlx/compile.cpp)。
使用 shapeless 编译的注意事项
请谨慎使用 shapeless 编译。由于形状变化不会触发重新编译,任何依赖输入形状的条件分支图都不会按预期工作。形状相关的计算很常见,而且有时很隐蔽,例如:
def fun(x): return x.reshape(x.shape[0] * x.shape[1], -1) compiled_fun = mx.compile(fun, shapeless=True) x = mx.random.uniform(shape=(2, 3, 4)) out = compiled_fun(x) x = mx.random.uniform(shape=(5, 5, 3)) # 报错:无法将 (5, 5, 3) 变形为 (6, -1) out = compiled_fun(x)第二次调用失败的原因是:reshape使用了第一次调用时x的静态形状(2 * 3 = 6),而 shapeless 模式下图不会按新形状重新构建。解决办法是改用flatten,避免在图中硬编码形状:
def fun(x): return x.flatten(0, 1) compiled_fun = mx.compile(fun, shapeless=True) x = mx.random.uniform(shape=(2, 3, 4)) out = compiled_fun(x) x = mx.random.uniform(shape=(5, 5, 3)) # 正常 out = compiled_fun(x)另外需要留意:shapeless=True并非适用于所有函数,尝试编译不支持 shapeless 的函数会抛错;并且即便启用shapeless,改变输入的维度数(ndim)或类型仍会触发重新编译(见 python/src/transforms.cpp 的签名文档)。shapeless的边界情况在 python/tests/test_compile.py 中有大量测试覆盖,包括与 broadcast、reduction、gather、full_like、量化矩阵乘等算子的组合,可作为判断"哪些图可以安全 shapeless 编译"的参考。
八、小结与最佳实践
结合 compile.rst 与源码实现,可以把mx.compile的使用要点归纳为:
- 编译可复用的函数:首次编译成本高,但结果会被缓存;避免在循环中反复创建编译函数(缓存键基于函数地址,见 mlx/compile.cpp)。
- 警惕重编译触发条件:输入 shape/ndim、dtype、输入个数的变化都会导致(部分或全部)重新编译。
- 保持函数纯净:编译函数内部不能对数组求值(打印会崩溃);有外部可变状态时,用
inputs=/outputs=捕获;训练时记得把mx.random.state一并捕获。 - 编译最外层:让
compile有机会融合尽可能多的算子(受max_compile_depth=11与max_compile_arrays=24约束)。 - 调试用开关:
mx.disable_compile()或环境变量MLX_DISABLE_COMPILE可全局关闭编译(编译模式在 mlx/compile.h 中定义为disabled / no_simplify / no_fuse / enabled四档),enable_compile()可重新开启并覆盖环境变量。 - shapeless 谨慎用:它跳过 shape 相关的重编译,但要求计算图本身不依赖静态形状(多用
flatten而非手写reshape维度)。 - 组合变换时显式编译:对编译函数再施加变换后默认不编译,需要编译时显式再包一层
mx.compile。
进一步深入时,可以参考 python/tests/test_compile.py 中的完整测试集(常量、无穷值、闭包捕获、kwargs、多线程编译、随机状态捕获等场景),以及 function_transforms.rst 中关于函数变换组合性的系统说明。
【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考