news 2026/9/11 15:45:02

从零实现线性回归:MLX 中 mx.grad 自动微分与 SGD 训练实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零实现线性回归:MLX 中 mx.grad 自动微分与 SGD 训练实战

从零实现线性回归: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。这样训练结束后可以直接度量ww_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

生成过程分三步:

  1. 采样设计矩阵Xmx.random.normal((num_examples, num_features))生成 1000×100 的标准正态随机矩阵,每行是一条样本;
  2. 采样真值参数w_starmx.random.normal((num_features,))生成 100 维标准正态向量;
  3. 合成带噪标签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 的自动微分作用于函数而非隐式计算图——没有backwardzero_gradrequires_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)也明确:gradvmap等变换可以按任意顺序、任意深度组合,例如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)

每一步发生了什么:

  1. grad = grad_fn(w):通过自动微分构建计算图,此刻并未真正计算
  2. w = w - lr * grad:数组赋值只是记录新的图节点;
  3. mx.eval(w):显式触发求值,真正调度 GPU/CPU 执行整条计算链。

初始参数缩放为1e-2是为了避免初始化幅值过大;w_star由标准正态采样生成,其典型范数约sqrt(100) = 10,因此1e-2的初始尺度是合理的较小起点。

为什么要调用 mx.eval:惰性求值模型

MLX 采用惰性求值:执行X @ wmx.square等操作时,实际不发生计算,只是把操作记录进一张计算图(compute graph)。只有当调用mx.evalprint.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.ndarraymemoryview访问以及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.00364
  • loss = loss_fn(w):训练后损失,理想情况下接近噪声方差量级(eps幅度1e-2,平方均值约为(1e-2)^2 / 2的量级);
  • error_norm = mx.sum(mx.square(w - w_star)).item() ** 0.5:学习到的w与真值w_starL2 距离。由于噪声的存在,它不会精确为 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_featuresnum_examples时直观评估性能。


七、延伸一步:从线性回归到逻辑回归(分类)

理解线性回归后,只需修改两处即可得到分类模型。仓库中的 examples/python/logistic_regression.py 展示了这一迁移:

  1. 标签变为二值y = (X @ w_star) > 0(正负例各半);
  2. 损失换为 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.cpp
  • matmul/mean/square等算子声明:mlx/ops.h

结语

通过这个最小可运行的线性回归例子,你实际上已经掌握了 MLX 三块核心拼图:数组与算子mx.random.normal@mx.meanmx.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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 15:44:07

STM32G0与MT6835的SPI通信陷阱:幽灵SCK与CS时序真相

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 15:42:51

微信小程序云开发直连数据库实战指南

1. 项目概述:小程序开发的新范式去年接手一个电商小程序项目时,客户要求两周内上线MVP版本。传统开发方式需要同时搭建后端服务、设计API接口、开发管理后台,时间根本来不及。最终我们采用微信云开发直连数据库的方案,前端团队独立…

作者头像 李华
网站建设 2026/9/11 15:42:13

C语言数组深度解析:从内存布局到越界调试

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 15:41:43

qwen-code 结构化调试方法论:用假设驱动循环替代盲目修复

qwen-code 结构化调试方法论:用假设驱动循环替代盲目修复 【免费下载链接】qwen-code An open-source AI coding agent that lives in your terminal. 项目地址: https://gitcode.com/GitHub_Trending/qw/qwen-code 导读 qwen-code(项目仓库&…

作者头像 李华