ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

MLX 实战:用 mlx.nn 从零构建多层感知机(MLP)完成 MNIST 分类

MLX 实战:用 mlx.nn 从零构建多层感知机(MLP)完成 MNIST 分类 MLX 实战用 mlx.nn 从零构建多层感知机MLP完成 MNIST 分类【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx本文是 MLXApple 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 的三个核心模块以及用于数据预处理的 NumPyimport 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_layers2, input_dim784, hidden_dim32, output_dim10时得到[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还支持biasFalse关闭偏置以及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.pylogits未归一化的模型输出targets可以是类别索引此时形状为 logits 去掉axis维也可以是各类别概率/one-hot 向量形状与 logits 一致axissoftmax 作用的轴默认-1label_smoothing标签平滑因子取值[0, 1)默认0reductionnone | mean | sum默认none示例中通过外层mx.mean完成mean归约。评估函数则直接比较预测类别与真实标签def eval_fn(model, X, y): return mx.mean(mx.argmax(model(X), axis1) y)mx.argmax沿类别轴axis1取出每个样本预测的类别索引与标签数组做布尔比较再取平均即得到准确率。四、设置超参数并加载 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_layers2隐藏层数量本例为 2 个 32 维隐藏层hidden_dim32每个隐藏层的神经元数num_classes10输出类别数MNIST 数字 0–9batch_size256每个 mini-batch 的样本数num_epochs10对整个训练集遍历的次数learning_rate1e-1SGD 学习率数据加载依赖官方 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 models trainable parameters loss_and_grad_fn nn.value_and_grad(model, loss_fn) # Instantiate the optimizer optimizer optim.SGD(learning_ratelearning_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(fEpoch {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_gradmx.value_and_grad通用函数变换不是同一个东西前者专门处理模型参数的递归结构后者是 MLX 函数变换原语。在训练 MLX 模型时应使用nn.value_and_grad。3. SGD 优化器一步完成状态与参数更新optimizer optim.SGD(learning_ratelearning_rate) optimizer.update(model, grads)optim.SGD的实现位于 python/mlx/optimizers/optimizers.py其更新规则为v_{t1} μv_t (1-τ)g_tw_{t1} w_t - λv_{t1}其中μ为动量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),仅供参考
返回列表