MLX 实战:用 mlx.nn 从零构建多层感知机(MLP)完成 MNIST 分类
【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx
本文是 MLX(Apple silicon 上的 NumPy 风格数组框架)官方示例教程《Multi-Layer Perceptron》的完整实践指南,对应仓库文档 docs/src/examples/mlp.rst。教程以 MNIST 手写数字分类为目标任务,演示如何使用mlx.nn定义模型、使用mlx.optimizers执行 SGD 优化,并通过mlx.core完成张量运算,最终在苹果芯片上用约 10 个 epoch 把测试准确率训练到约 95%。读完本文,你将掌握 MLX 构建与训练神经网络的标准工作流:继承nn.Module定义网络、用nn.value_and_grad自动求梯度、用optimizer.update(model, grads)一步完成参数更新,并理解 MLX 惰性求值(lazy evaluation)下mx.eval的作用。
一、导入 MLX 核心包
训练脚本的第一步是导入 MLX 的三个核心模块,以及用于数据预处理的 NumPy:
import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim import numpy as npmlx.core(别名mx)提供数组、自动微分(mx.value_and_grad)、随机数(mx.random)等基础能力,对应 mlx/core 层源码;mlx.nn(别名nn)提供Module基类、Linear等层实现、损失函数与value_and_grad等训练辅助函数,对应 python/mlx/nn/;mlx.optimizers(别名optim)提供SGD、Adam等优化器实现,对应 python/mlx/optimizers/optimizers.py。
二、用 nn.Module 定义 MLP 模型
MLX 中自定义网络的标准做法是继承mlx.nn.Module,其背后有一套完整的参数注册机制,核心定义在 python/mlx/nn/layers/base.py 的Module类中。构造新模块遵循两步惯例:
- 在
__init__中设置参数和/或子模块; - 在
__call__中实现前向计算。
class MLP(nn.Module): def __init__( self, num_layers: int, input_dim: int, hidden_dim: int, output_dim: int ): super().__init__() layer_sizes = [input_dim] + [hidden_dim] * num_layers + [output_dim] self.layers = [ nn.Linear(idim, odim) for idim, odim in zip(layer_sizes[:-1], layer_sizes[1:]) ] def __call__(self, x): for l in self.layers[:-1]: x = mx.maximum(l(x), 0.0) return self.layers-1逐行解读关键点:
super().__init__()是必须的,它初始化Module内部的参数字典、冻结集合(_no_grad)与训练状态标志(见 base.py 的 Module.init)。layer_sizes构造出维度链:[input_dim, hidden_dim, hidden_dim, ..., output_dim],例如num_layers=2, input_dim=784, hidden_dim=32, output_dim=10时得到[784, 32, 32, 10]。- 将多个
nn.Linear放进一个 Python 列表,Module的递归机制(parameters()/filter_and_map)会自动把它们登记为子模块并收集其中所有mx.array参数,因此不需要像某些框架那样使用特殊的"模块列表"容器。 - 前向传播对除最后一层外的所有线性层施加 ReLU 激活(
mx.maximum(x, 0.0)),最后一层直接输出 logits,不做 softmax——交叉熵损失内部会处理。
nn.Linear是 MLX 中最基础的层,源码位于 python/mlx/nn/layers/linear.py,其数学形式为y = xW^T + b:权重W形状为[output_dims, input_dims],偏置b形状为[output_dims]。参数初始化采用均匀分布U(-k, k),其中k = 1/sqrt(input_dims);前向计算在带偏置时调用mx.addmm(融合的矩阵乘加),无偏置时直接使用矩阵乘法x @ W.T。Linear还支持bias=False关闭偏置,以及to_quantized()方法将层转换为量化版本(QuantizedLinear/QQLinear),可用于后续推理压缩。
三、定义损失函数与评估函数
损失函数对每个样本的交叉熵取平均。mlx.nn.losses子包提供了若干常用损失函数的实现:
def loss_fn(model, X, y): return mx.mean(nn.losses.cross_entropy(model(X), y))nn.losses.cross_entropy的完整签名与语义定义在 python/mlx/nn/losses.py:
logits:未归一化的模型输出;targets:可以是类别索引(此时形状为 logits 去掉axis维),也可以是各类别概率/one-hot 向量(形状与 logits 一致);axis:softmax 作用的轴,默认-1;label_smoothing:标签平滑因子,取值[0, 1),默认0;reduction:'none' | 'mean' | 'sum',默认'none',示例中通过外层mx.mean完成mean归约。
评估函数则直接比较预测类别与真实标签:
def eval_fn(model, X, y): return mx.mean(mx.argmax(model(X), axis=1) == y)mx.argmax沿类别轴(axis=1)取出每个样本预测的类别索引,与标签数组做布尔比较,再取平均即得到准确率。
四、设置超参数并加载 MNIST 数据
num_layers = 2 hidden_dim = 32 num_classes = 10 batch_size = 256 num_epochs = 10 learning_rate = 1e-1 # Load the data import mnist train_images, train_labels, test_images, test_labels = map( mx.array, mnist.mnist() )参数速查表:
| 参数 | 取值 | 含义 |
|---|---|---|
num_layers | 2 | 隐藏层数量(本例为 2 个 32 维隐藏层) |
hidden_dim | 32 | 每个隐藏层的神经元数 |
num_classes | 10 | 输出类别数(MNIST 数字 0–9) |
batch_size | 256 | 每个 mini-batch 的样本数 |
num_epochs | 10 | 对整个训练集遍历的次数 |
learning_rate | 1e-1 | SGD 学习率 |
数据加载依赖官方 mlx-examples 仓库提供的mnist数据加载器(mnist.mnist()),该 loader 不在本仓库内,需要从 mlx-examples 的mnist示例中获取并放置于脚本同目录后以import mnist方式引入。它返回训练集/测试集的图像与标签四个部分,示例中通过map(mx.array, ...)将 NumPy 数组统一转换为mx.array,此后所有计算都在 MLX 张量上进行。
五、构造 mini-batch 迭代器
由于选用 SGD(随机梯度下降),需要一个对训练集打乱顺序并切分为 mini-batch 的迭代器:
def batch_iterate(batch_size, X, y): perm = mx.array(np.random.permutation(y.size)) for s in range(0, y.size, batch_size): ids = perm[s : s + batch_size] yield X[ids], y[ids]np.random.permutation(y.size)生成一个打乱的索引排列,再转为mx.array;- 按
batch_size步长切片,X[ids]与y[ids]是 MLX 的索引操作(高级索引),返回对应 batch 的图像与标签; - 每个 epoch 重新调用该迭代器即可获得不同的随机顺序,实现每轮数据洗牌。
六、训练循环:value_and_grad、SGD 与 mx.eval
将以上所有部分组装成完整训练循环:
# Load the model model = MLP(num_layers, train_images.shape[-1], hidden_dim, num_classes) mx.eval(model.parameters()) # Get a function which gives the loss and gradient of the # loss with respect to the model's trainable parameters loss_and_grad_fn = nn.value_and_grad(model, loss_fn) # Instantiate the optimizer optimizer = optim.SGD(learning_rate=learning_rate) for e in range(num_epochs): for X, y in batch_iterate(batch_size, train_images, train_labels): loss, grads = loss_and_grad_fn(model, X, y) # Update the optimizer state and model parameters # in a single call optimizer.update(model, grads) # Force a graph evaluation mx.eval(model.parameters(), optimizer.state) accuracy = eval_fn(model, test_images, test_labels) print(f"Epoch {e}: Test accuracy {accuracy.item():.3f}")1. 实例化模型并强制求值
MLP(...)构造时参数尚处于惰性状态(MLX 的计算默认是惰性的,数组只在需要时物化)。mx.eval(model.parameters())会真正分配内存并完成参数初始化(权重为mx.random.uniform生成),这也是 MLX 惰性求值模型下创建模型后的标准一步,可参考 base.py 的类文档示例。
2. nn.value_and_grad 一次性返回损失与梯度
loss_and_grad_fn = nn.value_and_grad(model, loss_fn)nn.value_and_grad是一个针对模块的训练辅助函数,其实现见 python/mlx/nn/utils.py:它把loss_fn包装为对模型可训练参数(model.trainable_parameters())求导的函数,内部调用mx.value_and_grad,返回"损失值 + 关于所有可训练参数的梯度树"。
注意:nn.value_and_grad(针对模型的便捷封装)与mlx.core.value_and_grad(mx.value_and_grad,通用函数变换)不是同一个东西,前者专门处理模型参数的递归结构,后者是 MLX 函数变换原语。在训练 MLX 模型时应使用nn.value_and_grad。
3. SGD 优化器一步完成状态与参数更新
optimizer = optim.SGD(learning_rate=learning_rate) optimizer.update(model, grads)optim.SGD的实现位于 python/mlx/optimizers/optimizers.py,其更新规则为:v_{t+1} = μv_t + (1-τ)g_t,w_{t+1} = w_t - λv_{t+1},其中μ为动量(momentum,默认 0)、τ为权重衰减、λ为学习率。optimizer.update(model, grads)是Optimizer基类提供的方法(见同文件update定义),它会同时更新优化器内部状态(如动量缓冲)和模型参数,并自动推进步数计数——这正是"在单个调用中完成优化器状态与模型参数更新"的含义。
4. mx.eval 强制图求值
mx.eval(model.parameters(), optimizer.state)由于 MLX 是惰性求值的,训练循环中累积的计算图并不会立即执行;每个 batch 调用一次mx.eval会强制评估模型参数与优化器状态,使得训练真正推进。这是 MLX 训练循环与 PyTorch 等急切执行框架最显著的区别之一,也是官方 惰性求值指南 中强调的核心概念。
5. 每个 epoch 输出验证集准确率
eval_fn(model, test_images, test_labels)在完整测试集上计算准确率,accuracy.item()将 MLX 标量数组转换为 Python 浮点数以便格式化输出。MLX 的mx.eval同样保证了该评估计算被执行。
七、运行效果与进一步探索
按照上述配置(2 层隐藏层、每层 32 维、batch 256、10 个 epoch、学习率 0.1),模型在训练集上只需几次遍历即可达到约 95% 的测试准确率——对于 784 维输入、总参数约 2.7 万的浅层 MLP 而言,这是一个合理的基准表现。
本文演示的完整范式可以推广到更复杂的任务:将nn.Linear换成 python/mlx/nn/layers/ 中的卷积层、Transformer 层、Embedding 等,即可构建 CNN、Transformer 等架构;将optim.SGD换成optim.Adam等其他优化器(见 docs/src/python/optimizers/common_optimizers.rst),并配合 docs/src/python/nn/module.rst 中Module的freeze、train/eval、save_weights/load_weights等能力完成更完整的训练与部署流程。官方仓库还提供了线性回归(docs/src/examples/linear_regression.rst)、LLaMA 推理(docs/src/examples/llama-inference.rst)等更复杂的端到端示例,可作为下一步的参考。
【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考