从零训练321万参数小模型:MLX环境下Transformer训练链路拆解

不加载预训练权重,也不是 LoRA 或蒸馏。这个实验用 640 条中文短句,在 MLX 上从零训练一个约 321 万参数的字符级 Transformer,拆解了词表、层数、注意力头、损失函数和推理生成的完整链路。

在 Apple M3 Max、96 GB 统一内存与 macOS 26.5.2 的实测环境下,一位开发者没有加载 Qwen、DeepSeek 等预训练权重,也没有调用云端模型 API,而是用 MLX 0.32.2 从随机权重开始训练了一个约 321 万参数的字符级 Transformer 模型。训练数据只有 640 条中文短句,400 步更新,完整流程耗时约 7.319 秒。这个模型最终能从“小白在餐厅”续写出“餐厅。”这样的回答,但仍会把“紫色”答成“黄色”。这不是一次追求效果的模型竞赛,而是一次把 tokenizer、层数、注意力头、损失函数、权重保存与推理生成逐项拆开看的教学实验。

不是微调,也不是蒸馏:从随机权重开始的字符级模型

这次实验的关键前提是“从零训练”。所谓从零,并不是手写 GPU 驱动、矩阵乘法或自动求导,而是不继承任何已有模型的语言能力。模型结构由实验者自行指定,矩阵权重随机初始化,部分偏置和归一化参数按框架默认方式初始化,底层计算交给 MLX 完成。

实验没有使用 MLX-LM 的微调入口,而是直接使用 MLX 的神经网络组件搭建模型并训练。模型配置为 4 层 Transformer,每层 4 个注意力头,隐藏维度 256,FFN 中间宽度 1024,词表大小 112,上下文长度 128,总参数量约 321 万。

  • 不使用预训练权重,也不是 LoRA 或蒸馏。
  • 训练对象是字符级语言模型,一个字符对应一个编号,标点和换行也算字符。
  • 词表只从训练集生成,编号只是索引,不代表权重大小或重要性。

例如,在本次保存的词表里,“餐”对应编号 103。这个编号并不表示“餐”的权重是 103,也不意味着编号越大的字越重要。训练时,模型会把同一句话错开一个位置:输入“小 白 在 餐 厅”,目标则是“白 在 餐 厅 。”。也就是说,看到“小”预测“白”,看到“小白”预测“在”,看到“小白在”预测“餐”。不足 128 个位置的部分会被补齐,并让这些位置不参与 loss 计算。

层、头与 FFN 如何落到代码里

实验中最容易混淆的是“层”和“头”。整体结构是 4 层 Transformer,每层内部包含 4 个注意力头。层与层前后串联,每层里的多个头并行计算,因此并不是“4 × 4 = 16 层”。每个注意力头通过各自的投影权重处理输入,得到自己的 64 维表示。头学到什么并没有人为分工,不存在某个头固定负责“颜色”、另一个头固定负责“地点”的设定。

FFN 则对每个位置的表示分别加工。注意力负责在 token 之间交换信息,FFN 本身不直接进行跨 token 的注意力计算。代码层面,模型结构可以简化为词嵌入、位置编码、Transformer 编码器和输出线性层四部分:

  • nn.Embedding(112, 256):把 token 编号映射为 256 维向量。
  • nn.SinusoidalPositionalEncoding(256):告诉模型字符的先后位置。
  • nn.TransformerEncoder(num_layers=4, dims=256, num_heads=4, mlp_dims=1024, dropout=0.0, norm_first=True):4 层、每层 4 头、FFN 中间宽度 1024。
  • nn.Linear(256, 112, bias=False):对词表中的 112 个 token 分别打分。

数据真正进入模型时,核心流程是先做词嵌入和位置编码,再加上因果掩码,经过 4 层 Transformer,最后输出各 token 的原始分数。这里的 mask 很关键,它禁止当前位置看到后面的字符,防止模型在训练时“偷看”正确答案。虽然框架组件名为 TransformerEncoder,但实验显式加入因果掩码,将其用于自回归语言建模,并没有再额外接一个带交叉注意力的解码器。

训练循环、loss 与 400 步到底意味着什么

训练数据是在代码中组合人名、食物、地点和颜色生成的 640 条短句,例如“故事:小白在餐厅。 问:小白在哪里? 答:餐厅。”以及“故事:花朵是紫色的。 问:花朵是什么颜色? 答:紫色。”。这些数据被分为 576 条训练样本和 64 条验证样本,先按完整样本去重,再固定随机种子打乱、划分,两个集合不包含完全相同的样本。但两组资料仍然共用同一套句式,因此验证只能观察相似句式上的表现,不能证明模型已经理解通用中文。

损失函数采用交叉熵,比较模型给出的预测分数与正确 token。这里不是拿“预测编号 103”和“正确编号 31”做减法,而是关注模型给正确 token 分配的预测概率有多大。实验对故事、问题、答案中的所有有效“下一个字符”位置都计算损失,而不是只在答案区域计算,这也会影响后续对 loss 的理解。

优化器为 AdamW,learning_rate 为 0.001,weight_decay 为 0.01。每一步抽取 16 条样本,前向计算与反向传播得到 loss 和梯度,再限制过大的梯度,然后更新权重。训练设置中的 400 步,是更新权重 400 次,不是 400 层,也不是只处理或生成 400 个 token。

从结果看,训练集 loss 从 4.7344 降到 0.4677。数据准备、400 步训练、验证、保存和重载检查合计约 7.319 秒,不包含环境安装、代码编写、单测和事后复核。这个时间来自特定硬件上的小规模实验,不能外推为“训练一个会聊天的模型只要 7 秒”。

在生成评测中,模型对 64 条验证题逐条生成答案,去除首尾空白后做精确匹配,答对 26 条,即 40.625%。这是额外的生成评测,不是 loss 指标。loss 已经明显下降,但答案仍会出错,原因在于本实验的 loss 覆盖整个样本,模型学会常见句式、标点和固定措辞也能降低平均误差,却不等于它一定能正确提取每个问题需要的信息。

训练之后如何推理:权重保存、加载与贪心生成

训练结束后,实验保存了权重、词表和配置文件。仅权重文件约 12.86 MB,整个正式运行目录约 12.9 MB。代码实测的 MLX 内存峰值约 513 MB,这个数字不能理解为整个系统或 Python 进程的总内存,也不能把权重文件大小等同于运行内存。

权重、词表和配置需要配套加载。同一个编号在不同词表里可能对应不同字符,不能随意混用。重载时,实验创建全新模型实例,再从文件读取权重。固定输入的 logits 最大绝对差为 0,验证 loss 和生成文本也保持一致。

推理时,如果输入“故事:小白在餐厅。 问:小白在哪里? 答:”,这段文本共有 21 个字符 token。流程是:21 个 token 编号变成 21 × 256 的向量表示,依次经过 4 层,取最后一个位置的输出分数,选中 token 103,也就是“餐”;再把“餐”追加进去,继续预测后面的字,最终续写为“餐厅。”,接着生成换行并停止。

本实验选择最高分 token,属于贪心生成,而不是每次随机抽一个字。测试过程中没有执行优化器更新,权重不会因为提问而改变。为了让实现容易阅读,这个版本没有使用 KV Cache,每生成一个字符都会重新计算窗口内的上下文,因此更适合教学演示,而不是高性能推理服务。

这类实验的价值,不在于用几百条短句训练出一个可靠问答模型,而在于把大模型常见的抽象概念落到可运行的代码里:词表如何生成,字符如何变成编号,层数和注意力头如何组织,因果掩码如何防止看见未来,loss 为什么下降,推理时又如何一个字符一个字符续写。对于想在本地 Apple 芯片上理解 Transformer 训练链路的开发者来说,这种小模型实验提供了一条可以直接复现和拆解的路径。

原创文章,作者:点点,如若转载,请注明出处:https://www.dian8dian.com/cong-ling-xun-lian-321-wan-can-shu-xiao-mo-xing-mlx-huan

Like (0)
点点的头像点点
Previous 20小时前
Next 7小时前

相关推荐