从零实现线性回归:MLX 中 mx.grad 自动微分与 SGD 训练实战
【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx
导读
本文基于 docs/src/examples/linear_regression.rst 教程,讲解如何在 MLX(Apple silicon 上的数组计算框架)中用约 30 行 Python 代码从零实现一个线性回归模型:合成带噪数据集、定义均方误差损失、用mx.grad自动求梯度,再以随机梯度下降(SGD)迭代优化参数。读完本文,你将掌握 MLX 的核心编程范式——函数变换(function transform)、惰性求值(lazy evaluation)与mx.eval的配合时机,并能独立迁移到逻辑回归、MLP 等更复杂的模型上。
一、准备工作:导入 mlx.core 并设定问题元数据
线性回归的目标是学习参数向量w,使得X @ w尽可能逼近标签y。在 MLX 中,所有数组操作都经由mlx.core(通常缩写为mx)完成,它提供了类似 NumPy 的接口,但计算被调度到 Apple silicon 的 GPU/CPU 上执行。
import mlx.core as mx num_features = 100 num_examples = 1_000 num_iters = 10_000 # iterations of SGD lr = 0.01 # learning rate for SGD这段代码定义了一个经典的最小二乘回归问题:
num_features = 100:每条样本的特征维度;num_examples = 1_000:样本数量(设计矩阵的行数);num_iters = 10_000:SGD 的迭代步数;lr = 0.01:SGD 的学习率。
对照仓库中完整的可运行示例 examples/python/linear_regression.py,这些参数完全一致,是验证收敛行为的稳定起点。
说明:
mx中的数组默认类型为float32,这一点与 NumPy 不同(NumPy 默认float64)。本文示例不涉及类型显式转换,MLX 会根据输入自动推断。
二、合成数据集:设计矩阵、真值参数与高斯噪声
在真实数据不可得的情况下,教程采用"已知答案"的合成数据来检验学习算法:先随机生成一个"上帝视角"的真值参数w_star,再由它线性组合出带噪声的标签y。这样训练结束后可以直接度量w与w_star的距离,验证优化是否真正收敛。
# True parameters w_star = mx.random.normal((num_features,)) # Input examples (design matrix) X = mx.random.normal((num_examples, num_features)) # Noisy labels eps = 1e-2 * mx.random.normal((num_examples,)) y = X @ w_star + eps生成过程分三步:
- 采样设计矩阵
X:mx.random.normal((num_examples, num_features))生成 1000×100 的标准正态随机矩阵,每行是一条样本; - 采样真值参数
w_star:mx.random.normal((num_features,))生成 100 维标准正态向量; - 合成带噪标签
y:先计算无噪声响应X @ w_star(矩阵向量乘),再叠加幅度为1e-2的高斯噪声eps,模拟观测误差。
随机数背后的实现:Threefry PRNG
mx.random.normal并非朴素的伪随机实现。根据 docs/src/python/random.rst,MLX 遵循 JAX 的 PRNG 设计,采用可分裂(splittable)版本的Threefry计数器型伪随机数生成器。默认情况下所有采样函数使用隐式的全局 PRNG 状态;当需要精确复现或细粒度控制时,可以显式传入key:
key = mx.random.key(0) x = mx.random.normal((num_features,), key=key) # 同 key 得到同一序列矩阵乘@的实现位置
y = X @ w_star + eps中的@对应mx.matmul。在 mlx/ops.h 中可以看到其声明:MLX_API array matmul(const array& a, const array& b, StreamOrDevice s = {});。StreamOrDevice允许指定计算流或设备(CPU/GPU),默认交给框架自动调度。
三、损失函数与 mx.grad:函数式自动微分
有了数据和标签,接下来定义损失并求梯度。MLX 的自动微分作用于函数而非隐式计算图——没有backward、zero_grad、requires_grad这类 PyTorch 式 API,你只需要把"从参数到损失"的纯函数交给mx.grad,它会返回一个求梯度的新函数:
def loss_fn(w): return 0.5 * mx.mean(mx.square(X @ w - y)) grad_fn = mx.grad(loss_fn)这里的损失函数是半均方误差:
loss(w) = 0.5 * mean( (X @ w - y) ** 2 )mx.square逐元素平方(声明见 mlx/ops.h);mx.mean对全部元素求均值(声明见 mlx/ops.h);- 系数
0.5使梯度表达式更整洁,不改变最优解。
mx.grad(loss_fn)返回一个新函数grad_fn,默认对第一个参数求梯度。grad_fn(w)返回的是一个与w形状相同的梯度数组,而不是把梯度"存"在某个张量上——这正是函数式自动微分的核心:一切以函数的输入输出为边界。
函数变换可以任意组合
mx.grad属于 MLX 的"可组合函数变换"(composable function transformations),其输出仍是普通函数,因此可以继续被变换。在 docs/src/usage/function_transforms.rst 中给出了直观例子:
>>> mx.grad(mx.sin)(mx.array(mx.pi)) array(-1, dtype=float32) # 恰为 cos(pi) >>> mx.grad(mx.grad(mx.sin))(mx.array(mx.pi / 2)) array(-1, dtype=float32) # 恰为 -sin(pi/2)即grad(grad(fn))可以不断求出更高阶导数。快速入门文档(docs/src/usage/quick_start.rst)也明确:grad、vmap等变换可以按任意顺序、任意深度组合,例如grad(vmap(grad(fn)))。
进阶:value_and_grad 与 argnums
如果既要损失值又要梯度,应使用mx.value_and_grad,避免前向传播被重复计算:
loss_and_grad_fn = mx.value_and_grad(loss_fn) loss, grad = loss_and_grad_fn(w)若要对非首个参数求梯度,可用argnums指定位置;梯度还能作用于任意嵌套的list/tuple/dict参数树,且梯度保持与参数相同的树形结构(详见 docs/src/usage/function_transforms.rst)。
四、SGD 训练循环:梯度下降与 mx.eval 的配合
初始化参数后,反复执行"求梯度 → 沿负梯度方向更新"即可:
w = 1e-2 * mx.random.normal((num_features,)) for _ in range(num_iters): grad = grad_fn(w) w = w - lr * grad mx.eval(w)每一步发生了什么:
grad = grad_fn(w):通过自动微分构建计算图,此刻并未真正计算;w = w - lr * grad:数组赋值只是记录新的图节点;mx.eval(w):显式触发求值,真正调度 GPU/CPU 执行整条计算链。
初始参数缩放为1e-2是为了避免初始化幅值过大;w_star由标准正态采样生成,其典型范数约sqrt(100) = 10,因此1e-2的初始尺度是合理的较小起点。
为什么要调用 mx.eval:惰性求值模型
MLX 采用惰性求值:执行X @ w、mx.square等操作时,实际不发生计算,只是把操作记录进一张计算图(compute graph)。只有当调用mx.eval、print、.item()或转成 NumPy 数组时,计算才真正发生。底层实现见 mlx/transforms.cpp:eval会检查输出数组中是否存在状态为unscheduled的节点,若有则触发eval_impl完成调度与执行,否则只做一次wait(等待已有结果)。因此对已求值数组重复调用mx.eval是安全且近乎零开销的。
在训练循环的每次迭代末尾调用mx.eval(w)是官方推荐的做法:docs/src/usage/lazy_evaluation.rst 指出,大多数数值计算都有迭代外层循环(如 SGD),"在每个外层循环迭代处调用eval是自然且通常高效的"。
eval 频率的权衡
- 过于频繁:每次求值都有固定开销,例如在循环内对中间变量逐个
mx.eval是不必要的; - 过于稀少:计算图无限增长,图的构建与内存占用会引入随图规模增长的轻微开销;
- 经验区间:单次求值覆盖几十到几千个操作都是合适的;
- 陷阱:把标量数组用于控制流(如
if y > 0:)会触发隐式求值,虽能工作但可能因求值过频而低效。
此外还有几条隐式求值规则:print数组、array.item()、转numpy.ndarray、memoryview访问以及mx.save都会自动触发求值(详见 docs/src/usage/lazy_evaluation.rst)。
五、验证收敛:损失与参数误差
训练结束后,用两个指标验证模型学得好不好:
loss = loss_fn(w) error_norm = mx.sum(mx.square(w - w_star)).item() ** 0.5 print( f"Loss {loss.item():.5f}, |w-w*| = {error_norm:.5f}, " ) # Should print something close to: Loss 0.00005, |w-w*| = 0.00364loss = loss_fn(w):训练后损失,理想情况下接近噪声方差量级(eps幅度1e-2,平方均值约为(1e-2)^2 / 2的量级);error_norm = mx.sum(mx.square(w - w_star)).item() ** 0.5:学习到的w与真值w_star的L2 距离。由于噪声的存在,它不会精确为 0,但应当非常小(示例输出约为0.00364);.item():把标量数组取成 Python 浮点数——这会触发一次求值,并顺带将结果打印出来。
六、完整可运行代码(含计时)
仓库中的 examples/python/linear_regression.py 是上述教程的完整版,额外用time.perf_counter()统计了迭代吞吐量。以下为完整实现,可直接复制运行:
import time import mlx.core as mx num_features = 100 num_examples = 1_000 num_iters = 10_000 lr = 0.01 # True parameters w_star = mx.random.normal((num_features,)) # Input examples (design matrix) X = mx.random.normal((num_examples, num_features)) # Noisy labels eps = 1e-2 * mx.random.normal((num_examples,)) y = X @ w_star + eps # Initialize random parameters w = 1e-2 * mx.random.normal((num_features,)) def loss_fn(w): return 0.5 * mx.mean(mx.square(X @ w - y)) grad_fn = mx.grad(loss_fn) tic = time.perf_counter() for _ in range(num_iters): grad = grad_fn(w) w = w - lr * grad mx.eval(w) toc = time.perf_counter() loss = loss_fn(w) error_norm = mx.sum(mx.square(w - w_star)).item() ** 0.5 throughput = num_iters / (toc - tic) print( f"Loss {loss.item():.5f}, L2 distance: |w-w*| = {error_norm:.5f}, " f"Throughput {throughput:.5f} (it/s)" )运行前提:本机为 Apple silicon(M 系列芯片),并已安装 MLX(pip install mlx)。throughput一栏给出每秒完成的 SGD 迭代数,便于在调整num_features、num_examples时直观评估性能。
七、延伸一步:从线性回归到逻辑回归(分类)
理解线性回归后,只需修改两处即可得到分类模型。仓库中的 examples/python/logistic_regression.py 展示了这一迁移:
- 标签变为二值:
y = (X @ w_star) > 0(正负例各半); - 损失换为 logistic 损失(数值稳定的对数损失写法):
def loss_fn(w): logits = X @ w return mx.mean(mx.logaddexp(0.0, logits) - y * logits) grad_fn = mx.grad(loss_fn)这里用mx.logaddexp(0.0, logits)实现log(1 + exp(logits)),避免指数上溢。学习率取lr = 0.1,训练结束后用准确率评估:
final_preds = (X @ w) > 0 acc = mx.mean(final_preds == y) print(f"Loss {loss.item():.5f}, Accuracy {acc.item():.5f} ...")可以看到,训练循环、mx.grad的使用方式、mx.eval的调度模式完全不变——变化的只是损失函数与标签定义。这正是函数式自动微分范式的复用价值:从回归到分类,再到 docs/src/examples/mlp.rst 中的多层感知机,核心骨架始终一致。
八、本文涉及的核心源码与文档索引
为便于继续深入,下面列出本文依据的主要仓库资源:
- 教程原文:docs/src/examples/linear_regression.rst
- 完整线性回归示例:examples/python/linear_regression.py
- 逻辑回归示例:examples/python/logistic_regression.py
- 惰性求值与 eval 时机:docs/src/usage/lazy_evaluation.rst
- 函数变换(grad/value_and_grad/vmap):docs/src/usage/function_transforms.rst
- 快速入门:docs/src/usage/quick_start.rst
- 随机数生成(Threefry PRNG 与 key):docs/src/python/random.rst
eval底层实现:mlx/transforms.cppmatmul/mean/square等算子声明:mlx/ops.h
结语
通过这个最小可运行的线性回归例子,你实际上已经掌握了 MLX 三块核心拼图:数组与算子(mx.random.normal、@、mx.mean、mx.square)、函数式自动微分(mx.grad及其可组合性)、惰性求值与显式求值(mx.eval的正确时机)。这三者共同构成了后续阅读 docs/src/examples/mlp.rst、docs/src/examples/llama-inference.rst 等进阶教程,以及编写自定义训练循环的基础。
【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考