新闻详情

新闻详情

首页 / 资讯中心 / 详情

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

发布时间:2026/9/11 15:45:03来源:尧图网络
从零实现线性回归: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 教程讲解如何在 MLXApple 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_000SGD 的迭代步数lr 0.01SGD 的学习率。对照仓库中完整的可运行示例 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生成过程分三步采样设计矩阵Xmx.random.normal((num_examples, num_features))生成 1000×100 的标准正态随机矩阵每行是一条样本采样真值参数w_starmx.random.normal((num_features,))生成 100 维标准正态向量合成带噪标签y先计算无噪声响应X w_star矩阵向量乘再叠加幅度为1e-2的高斯噪声eps模拟观测误差。随机数背后的实现Threefry PRNGmx.random.normal并非朴素的伪随机实现。根据 docs/src/python/random.rstMLX 遵循 JAX 的 PRNG 设计采用可分裂splittable版本的Threefry计数器型伪随机数生成器。默认情况下所有采样函数使用隐式的全局 PRNG 状态当需要精确复现或细粒度控制时可以显式传入keykey mx.random.key(0) x mx.random.normal((num_features,), keykey) # 同 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.hmx.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, dtypefloat32) # 恰为 cos(pi) mx.grad(mx.grad(mx.sin))(mx.array(mx.pi / 2)) array(-1, dtypefloat32) # 恰为 -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.cppeval会检查输出数组中是否存在状态为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( fLoss {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( fLoss {loss.item():.5f}, L2 distance: |w-w*| {error_norm:.5f}, fThroughput {throughput:.5f} (it/s) )运行前提本机为 Apple siliconM 系列芯片并已安装 MLXpip 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(fLoss {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/vmapdocs/src/usage/function_transforms.rst快速入门docs/src/usage/quick_start.rst随机数生成Threefry PRNG 与 keydocs/src/python/random.rsteval底层实现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),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

FlatBuffers Go 实战:基于 examples/go-echo 构建跨网络传输的零拷贝序列化示例 2026/9/11 16:27:11

FlatBuffers Go 实战:基于 examples/go-echo 构建跨网络传输的零拷贝序列化示例

FlatBuffers Go 实战:基于 examples/go-echo 构建跨网络传输的零拷贝序列化示例 【免费下载链接】flatbuffers FlatBuffers: Memory Efficient Serialization Library 项目地址: https://gitcode.com/GitHub_Trending/fl/flatbuffers 本篇指南以仓库 example…

阅读更多 →
GHelper:替代 Armoury Crate 的轻量方案,5 步调到位 2026/9/11 16:27:11

GHelper:替代 Armoury Crate 的轻量方案,5 步调到位

GHelper:替代 Armoury Crate 的轻量方案,5 步调到位 【免费下载链接】g-helper Lightweight Armoury Crate alternative for Asus laptops with nearly the same functionality. Works with ROG Zephyrus, Flow, TUF, Strix, Scar, ProArt, Vivobook, Ze…

阅读更多 →
上位机开发实战:从通信协议到工业级应用的三层架构 2026/9/11 16:27:11

上位机开发实战:从通信协议到工业级应用的三层架构

1. 这不是“转行”,是技术栈的精准迁移:一个25届应届生的真实上位机突围路径 “考研失利转行上位机,一周拿2个offer”——这个标题乍看像爽文,但在我带过的37个应届生项目里,它背后藏着一条被严重低估的、极其务实的技…

阅读更多 →
Material for MkDocs 教程体系:从博客搭建到社交卡片定制的完整实战路径 2026/9/11 16:27:11

Material for MkDocs 教程体系:从博客搭建到社交卡片定制的完整实战路径

Material for MkDocs 教程体系:从博客搭建到社交卡片定制的完整实战路径 【免费下载链接】mkdocs-material Documentation that simply works 项目地址: https://gitcode.com/GitHub_Trending/mk/mkdocs-material Material for MkDocs 在官方文档中专门设立了…

阅读更多 →
基于SpringBoot和MD5去重的校园网盘系统设计 2026/9/11 16:27:11

基于SpringBoot和MD5去重的校园网盘系统设计

简介:这是一份基于SpringBoot的校园网盘系统毕业设计源码与数据库资源,采用B/S架构,前端结合HTML、CSS、JavaScript、jQuery与Bootstrap,后端使用SpringBoot,配合MySQL数据库与Tomcat部署,可直接导入运行。…

阅读更多 →
IWOA-BiLSTM:改进鲸鱼算法优化双向LSTM超参 2026/9/11 16:24:11

IWOA-BiLSTM:改进鲸鱼算法优化双向LSTM超参

简介:本资源是一套面向高校科研人员与算法工程师的MATLAB时间序列预测实践代码包,聚焦于改进型鲸鱼优化算法(IWOA)与双向长短期记忆网络(BiLSTM)的融合建模与性能对比。资源解决了传统BiLSTM超参数调优依赖…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞