)
MLX 快速上手指南从数组创建到可组合函数变换基于 mlx.core 实战【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlxMLX 是面向 Apple silicon 的机器学习数组框架其 Python API 高度贴近 NumPy同时提供可组合的函数变换自动微分、自动向量化与惰性计算机制。本篇基于官方 Quick Start 文档带你完成从创建mlx.core.array、理解惰性求值到使用grad/vmap/vjp/jvp/value_and_grad构建可微分计算链路的完整入门读完即可动手编写并调试你自己的 MLX 训练脚本。安装与导入MLX 的 Python 包发布在 PyPI 上在 macOS 上直接安装pip install mlx在 Linux 上可以安装带 CUDA 后端的版本或仅使用 CPU 的版本pip install mlx[cuda] # Linux CUDA 后端 pip install mlx[cpu] # Linux CPU-only导入并开始使用import mlx.core as mxmlx.core通常简写为mx是 MLX 的核心模块提供数组类型、运算算子与函数变换更上层的神经网络模块mlx.nn与优化器mlx.optimizers则位于 Python 侧见 python/mlx 目录。创建数组Array 基础MLX 的核心数据类型是array对应 C 层的mlx::array见 mlx/array.h。创建一个数组最简单的方式是直接传入 Python 列表 import mlx.core as mx a mx.array([1, 2, 3, 4]) a.shape (4,) a.dtype int32 b mx.array([1.0, 2.0, 3.0, 4.0]) b.dtype float32两个要点a.shape返回数组各维度的长度元组(4,)表示这是一维、长度为 4 的数组a.dtype返回元素类型。整型列表默认推断为int32浮点列表默认推断为float32——这与 NumPy 的默认行为不同NumPy 默认int64/float64在跨框架移植代码时需要留意。如果你需要显式指定类型可以在创建时传入dtype参数例如mx.array([1, 2, 3], dtypemx.float16)。MLX 支持的数据类型还包括bf16、float16、int8、uint8、complex64等具体以 mlx/dtype.h 中的定义为准。惰性求值操作并不立即计算MLX 的核心设计之一是惰性计算lazy evaluation所有运算只是把操作记录到一张计算图compute graph上并不会真正执行计算直到结果被“需要”时才进行求值。这与 PyTorch 的立即执行eager风格有本质区别其设计动机与权衡详见 惰性求值指南。 c a b # c 尚未被求值 mx.eval(c) # 显式求值 c c a b print(c) # 打印也会触发求值 array([2, 4, 6, 8], dtypefloat32) c a b import numpy as np np.array(c) # 转成 numpy 数组同样会触发求值 array([2., 4., 6., 8.], dtypefloat32)在上面的例子中a b只是构造了计算图节点mx.eval(c)才真正执行加法并落盘结果。何时会“自动求值”除了显式调用mx.evalMLX 在若干场景下会自动对数组求值打印数组print(c)通过np.array(c)将数组转换为numpy.ndarray对标量数组调用array.item()例如在训练循环里把 loss 追加到列表时通过memoryview访问数组内存使用mx.save等保存函数持久化数组见 mlx/io/load.cpp。底层实现上Python 层的mx.eval由 python/src/transforms.cpp 绑定它会先用tree_flatten把传入的任意参数单个数组或 list/tuple/dict 构成的数组树展平再在释放 GIL 的情况下调用 C 层的eval因此你可以一次求值一整棵参数树例如mx.eval(loss, model.parameters())。求值时机与性能权衡惰性求值带来两个直接好处只计算你用到的输出。如果一个函数返回多个值而你只取其中一个其余分支的昂贵计算不会真正执行但计算图仍会被构建这部分开销依然存在降低峰值内存。例如实例化一个很大的模型时权重初始化并不会立刻占用内存如果你随后用 float16 的权重替换如model.load_weights(weights_fp16.safetensors)峰值内存可显著低于立即执行eager方案。至于何时调用mx.eval核心权衡是图太小则每次求值的固定开销占比过高图太大则单次求值成本随图规模增长。官方建议单次求值覆盖“几十到几千个操作”的图规模都比较合适对于 SGD 这类带迭代外循环的训练最自然的方式是在每个 batch 迭代结束时统一求值一次例如for batch in dataset: # 此刻尚未发生任何求值 loss, grads value_and_grad_fn(model, batch) # 仍然没有求值 optimizer.update(model, grads) # 在这里统一求值 loss 与新参数一次跑完前向、反向与优化器更新 mx.eval(loss, model.parameters())一个需要特别注意的坑用标量数组做控制流会触发隐式求值例如if y 0:这种写法会在条件判断处求值整个图。虽然它能工作、甚至能与梯度变换配合但如果求值发生得过于频繁性能会明显劣化使用需谨慎。另外对同一批数组重复调用mx.eval是安全的相当于空操作no-op。函数变换grad、vmap 及其组合MLX 提供标准的函数变换function transformationsgrad自动微分与vmap自动向量化并支持任意顺序、任意深度的组合例如grad(vmap(grad(fn)))完全合法。核心思想是每个变换都返回一个新的函数而新函数可以继续被变换。完整的 API 说明见 函数变换文档绑定实现在 python/src/transforms.cpp。grad一阶与高阶导数最简单的例子是对标量函数求导 x mx.array(0.0) mx.sin(x) array(0, dtypefloat32) mx.grad(mx.sin)(x) array(1, dtypefloat32) mx.grad(mx.grad(mx.sin))(x) array(-0, dtypefloat32)mx.grad(mx.sin)返回的正是 sin 的导数函数 cos在 0 处值为 1mx.grad(mx.grad(mx.sin))则是二阶导-sin在 0 处为 -0。对grad的输出再套grad永远可行你可以一直得到更高阶导数。grad默认对函数的第一个参数求导也可以用argnums指定用argnames指定具名参数详见 python/src/transforms.cpp 中的绑定签名def loss_fn(w, x, y): return mx.mean(mx.square(w * x - y)) w mx.array(1.0) x mx.array([0.5, -0.5]) y mx.array([1.5, -1.5]) grad_fn mx.grad(loss_fn) # 对 w 求导 print(grad_fn(w, x, y)) # array(-1, dtypefloat32) grad_fn mx.grad(loss_fn, argnums1) # 改为对 x 求导 print(grad_fn(w, x, y)) # array([-1, 1], dtypefloat32)此外grad还支持对任意嵌套的 Python 容器list、tuple、dict求梯度梯度会保持与输入相同的树结构。例如参数以字典形式组织时params {weight: mx.array(1.0), bias: mx.array(0.0)} grads mx.grad(loss_fn)(params, x, y) # {weight: array(-1, dtypefloat32), bias: array(0, dtypefloat32)}如果你来自 PyTorch注意在 MLX 中你不再需要backward、zero_grad、detach或requires_grad这类机制——梯度变换直接作用于函数本身。若需要阻断某条路径的梯度传播请使用mx.stop_gradient。value_and_grad同时拿到函数值与梯度分别调用loss_fn和grad_fn会带来大量重复计算。mx.value_and_grad一次调用同时返回函数值与梯度是训练循环中的首选 loss_and_grad_fn mx.value_and_grad(loss_fn) loss, dloss_dw loss_and_grad_fn(w, x, y) print(loss) # array(1, dtypefloat32) print(dloss_dw) # array(-1, dtypefloat32)其 C 实现python/src/transforms.cpp支持argnums整数或序列与argnames两种参数选择方式grad本质上是对value_and_grad的封装只保留梯度部分。vjp 与 jvp向量-雅可比积除grad外MLX 还提供两种更底层的微分变换vjp(fun, primals, cotangents)向量-雅可比积vector-Jacobian product输入反向模式微分中的“上游梯度”返回函数输出与该上游梯度的乘积是反向传播的基石jvp(fun, primals, tangents)雅可比-向量积Jacobian-vector product输入切向量返回沿该方向的方向导数。两者的 Python 绑定签名见 python/src/transforms.cpp测试用例覆盖了大量算子gather、scatter、slice_update 等的 vjp/jvp 正确性例如 python/tests/test_autograd.py 中的基本用法out, dout mx.vjp(fun, [mx.array(1.0)], [mx.array(2.0)])vmap自动向量化vmap把针对单样本编写的函数自动向量化从而省去手写批量循环。通过in_axes指定对每个输入在哪一个维度上向量化用out_axes指定输出中向量化轴的位置。典型对比示例完整代码见 函数变换文档xs mx.random.uniform(shape(4096, 100)) ys mx.random.uniform(shape(100, 4096)) def naive_add(xs, ys): return [xs[i] ys[:, i] for i in range(xs.shape[0])] # 对 x 的第 2 维、y 的第 1 维做向量化 vmap_add mx.vmap(lambda x, y: x y, in_axes(0, 1))注意vmap并非对所有算子都可用——如果遇到ValueError: Primitives vmap not implemented.之类的报错说明该算子的 vmap 变换尚未实现。另外若需要为自定义函数提供手写的 vjp/jvp/vmap 规则可以使用mx.custom_function装饰器见 python/src/transforms.cpp。更进一步本指南覆盖了 Quick Start 的全部内容数组创建与类型推断、惰性求值机制、以及可任意组合的函数变换。继续深入的方向深入理解惰性求值的动机、隐式求值触发条件与图规模权衡惰性求值指南完整的函数变换 APIgrad、vmap、vjp、jvp、value_and_grad、compile等函数变换文档用mlx.nn与mlx.optimizers搭建真实模型线性回归示例、逻辑回归示例、MLP 文档阅读官方测试 python/tests/test_autograd.py 与 python/tests/test_vmap.py了解各类算子对变换的支持情况。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考