ARTICLE DETAIL

资讯详情

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

用MLX从零训练321万参数Transformer:完整实战记录

用MLX从零训练321万参数Transformer:完整实战记录 标题里“321万”这个数字看着挺唬人对吧其实在 LLM 动不动几百亿参数的今天这点参数量连“玩具”都算不上。但我还是坚持把它当成一个正经项目来做原因很简单我想把 Transformer 的训练过程彻底看明白。光看原理文章你永远不知道 loss 曲线长什么样、梯度在哪个环节会炸、数据该切成什么样模型才学得动。用 MLX 从零写一个模型亲手跑起来这些就全懂了。这篇文章不是复述 Transformer 论文也不是堆概念而是记录我完整实现并训练一个 321 万参数小模型的全过程。数据准备、模型搭建、训练循环、调参陷阱、推理生成每一步我都会讲清楚“为什么这么做”和“我踩过什么坑”。适合已经入门过 Transformer 基础、想亲手跑通一次完整训练流程的人当然你要是完全没接触过我也会尽量把关键点说明白。1. 项目整体设计与思路拆解1.1 为什么用 MLX 而不是 PyTorch做这个项目之前我其实纠结过一阵子用 PyTorch 的话生态最熟教程最多报错一搜就能找到答案。但最终选了 MLX有几个很实际的原因。MLX 是苹果开源的机器学习框架底层针对 Apple Silicon 做了优化和 PyTorch 的 MPS 后端相比它对内存的管理方式完全不同。Mac 的统一内存Unified Memory可以让 CPU 和 GPU 访问同一块内存MLX 直接把这个优势用到了极致——训练和推理时不需要在设备之间搬数据张量在 CPU 和 GPU 之间切换几乎是无感的。而 PyTorch 的 MPS 在某些算子上的支持仍不完整一个小 op 不支持就得改写法很打断节奏。另外 MLX 的 API 非常接近 NumPy学过 NumPy 的人几乎零成本上手。它不像 PyTorch 有那么多层级封装很多操作就是直接对多维数组做运算。这种“朴素感”反而很适合学习你可以清楚看到每个张量的形状变化不会被框架自动封装的黑盒逻辑带偏。我从搭模型到跑通第一个 batch花的时间比预想中短很多。还有一个朴素的原因我现在的主力机器就是 MacBookMLX 在 Apple Silicon 上跑小规模实验非常顺手。如果你也刚好用 Mac想本地跑点小实验又不想折腾 CUDA 相关的环境MLX 是值得试的路线。1.2 为什么把参数规模卡在 321 万这个数字不是拍脑袋定的是反推出来的。我当时的约束条件有三个第一在 MacBook 上训练不能太慢最好几十分钟内能看到明显效果第二模型得是一个结构完整的 Transformer该有的组件一个不能少第三要能通过参数核算清晰讲解每个模块的“分量”。综合下来我选定了这样一组配置词表大小 512、模型维度 d_model256、6 层 Transformer block、前馈网络维度 d_ff480、最大序列长度 128。这组配置下参数总量是 3,217,920 左右约等于 321 万。具体怎么算出来的第三章我会列一张完整的参数核算表从 embedding 到每一层的 QKV 矩阵都算得明明白白。这个规模的模型有一个明显优势它对硬件的门槛很低。M1/M2 芯片的 MacBook 就能轻松跑起来不需要外置 GPU不需要租云服务器甚至连风扇狂转的情况都很少。但它又足够复杂足以演示训练过程中所有关键现象——loss 下降、过拟合、梯度异常、采样生成退化等等。我始终觉得理解大模型最好的方式就是在一个能掌控的小模型上把每个细节都折腾一遍。1.3 环境准备与依赖安装MLX 的安装非常简单核心就一个 Python 包。我用的是 conda 管理环境你可以按自己的习惯来但建议不要直接装到系统 Python 里免得污染全局环境。# 创建一个干净的虚拟环境 conda create -n mlx-train python3.10 -y conda activate mlx-train # 安装 MLX pip install mlx # 顺便装上后续会用到的工具库 pip install numpy matplotlib装完之后验证一下能不能正常调用import mlx.core as mx # 确认能创建数组并且能跑到 GPU 上 x mx.array([1, 2, 3]) print(x.device) # 正常会输出 cpu 或 gpu # 如果下面的输出能正常显示说明 MLX 已经可以用了 a mx.ones((4, 4)) b mx.eye(4) print((a b).shape)我实际测试过在 macOS 13 以上、Apple Silicon 芯片的机器上MLX 基本能做到开箱即用。如果你用的是 Intel Mac也能跑但性能会差一些。另外建议把 MLX 版本保持在 0.5 以上因为早期版本的 API 变化比较大很多接口到了新版才稳定下来。2. 数据准备让模型先有“字”2.1 数据选型为什么用 TinyShakespeare训练语言模型第一步是找数据。我选了一个非常经典的小数据集TinyShakespeare。这是 Karpathy 在相关教程里反复用过的数据集内容是莎士比亚的戏剧文本纯文本格式大小约 1MB字符量在 40 万左右。对跑通训练流程来说这个量级非常合适——它足够让模型学到一些语言规律又不会大到让你在数据处理上浪费大量时间。用这个数据集还有一个好处它是纯英文文本没有复杂的编码问题。虽然我用的是 byte-level 级别的分词方案但 TinyShakespeare 基本都是标准 ASCII 字符处理起来很干净。你如果手头有其他语料比如中文小说、代码、法律条文也可以替换进去训练流程是不变的变的只是词表和词表大小。补充一句数据规模不一定要大但一定要“干净”。我见过很多人第一步就在数据上翻车——文件编码错了、文本里有大量乱码或者 HTML 标签、换行符不统一这些都会直接影响训练效果。所以我建议拿到任何数据先做一轮简单的清洗统一换行符、去掉过长的空白序列、剔除异常字符。这个步骤看起来不起眼但能省掉后面排查问题的大量时间。2.2 手写一个 512 词表的 TokenizerTransformer 不能直接处理文本字符它处理的是整数 token。所以需要一个分词器Tokenizer把文本切成一串整数 ID。现在的工业级方案大多是 BPEByte Pair Encoding或者 SentencePiece但在这个项目里我不想引第三方分词库而是自己实现一个简化版的 BPE目标就是把词表定在 512。思路其实很简单先把整个训练语料按 UTF-8 编码拆成字节每个字节看成“初始 token”然后统计所有相邻 token 对的频率每次把出现频率最高的那一对合并成一个新 token重复这个过程直到词表达到目标大小。这个算法就是 BPE 的原始思路实现起来也就几十行代码。import collections def build_bpe_vocab(text, target_vocab_size512): # 先拆成字节列表 tokens list(text.encode(utf-8)) vocab {i: bytes([i]) for i in range(256)} while len(vocab) target_vocab_size: # 统计相邻对频率 pairs collections.Counter(zip(tokens, tokens[1:])) if not pairs: break (a, b), _ pairs.most_common(1)[0] new_id max(vocab.keys()) 1 # 合并所有相邻的 (a, b) merged [] i 0 while i len(tokens): if i len(tokens) - 1 and tokens[i] a and tokens[i1] b: merged.append(new_id) i 2 else: merged.append(tokens[i]) i 1 tokens merged vocab[new_id] vocab[a] vocab[b] return vocab, tokens核心逻辑就这么点。实际做的时候要注意初始字节是 0-255所以 BPE 新增的 token ID 从 256 开始编号。当词表达到 512 时实际上我们获得了 256 个初始字节 token 加上 256 个合并后的高频词片。训练语料被转成整数序列后就可以直接喂给模型了。注意这里用的是“字节级 BPE”的简化版和 GPT 系列的 tokenizer 原理类似只是我砍掉了正则预切分、特殊 token 等复杂环节。如果你要处理中文等多字节语言字节级方案也能工作因为中文字符在 UTF-8 下会变成一个字节序列BPE 会自动学习出高频字节组合效果也还行。2.3 切分序列构建 (input, target) 数据对语言模型的训练任务本质上是“预测下一个 token”。也就是说给定前面一段文本让模型预测下一个字符/词是什么。所以要把一长串 token 序列切成很多个固定长度的样本对(input, target)。我在项目里设定了 block_size128意思是每次模型只能看到 128 个 token 的上下文。切分逻辑是这样的假设有一串 token[t0, t1, t2, ..., tn]那么第一个样本的 input 就是[t0, t1, ..., t126]target 就是[t1, t2, ..., t127]——整体右移一位。第二个样本从t1开始以此类推。def get_batch(data, batch_size, block_size): # 随机选择 batch_size 个起始位置 ix mx.random.randint(0, len(data) - block_size, (batch_size,)) x mx.stack([data[i:iblock_size] for i in ix.tolist()]) y mx.stack([data[i1:iblock_size1] for i in ix.tolist()]) return x, y注意这里 target 的第 k 个位置对应 input 第 k 个位置的“下一个 token”。训练时 loss 的 mask 不需要额外处理因为每个位置都有对应的预测目标。我按照 9:1 的比例把数据切成了训练集和验证集。验证集不参与梯度更新只用来观察模型有没有过拟合。这一步很关键很多人只在训练集上盯 loss结果模型背数据背得很开心一换数据就原形毕露。3. Transformer 架构逐段拆解3.1 模型骨架Embedding 与位置编码模型搭起来其实不复杂但每一步都应该清楚它在做什么。首先是 Token Embedding它把 token ID 映射成一个 256 维的稠密向量。本质上这就是一个可学习的查找表512 个 token每个 token 对应一个 256 维向量。class TransformerBlock(nn.Module): def __init__(self, d_model, d_ff, n_heads, dropout): super().__init__() self.ln1 nn.LayerNorm(d_model) self.ln2 nn.LayerNorm(d_model) self.attn nn.MultiheadAttention(d_model, n_heads, dropoutdropout) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model), ) self.dropout nn.Dropout(dropout)Embedding 层在 MLX 里可以直接用nn.Embedding(num_embeddings512, embedding_dim256)。但我不建议直接初始化后就用最好给一个缩放系数。我参考 GPT 系列的做法把 embedding 的权重乘以sqrt(d_model)也就是 16 倍根号下 256 的算术平方根。这个缩放的作用是让输入到 Attention 的向量尺度保持稳定避免初始状态过小或过大影响梯度流动。位置编码我用的是可学习的方案初始化一个(128, 256)的参数矩阵把它和 token embedding 直接相加。为什么需要位置编码因为 Attention 本身是“无序”的它只看 token 之间的相关性不知道谁在前谁在后。如果不加位置信息模型会把“我爱你”和“你爱我”的表示混成一样的。可学习位置编码的好处是简单、直观模型在训练中自己学会如何使用位置信息。当然你也可以换成 RoPE旋转位置编码更现代但可学习方案对这个项目完全够用。3.2 自注意力QKV 的计算过程自注意力是 Transformer 的核心也是最容易让人看得一头雾水的地方。我建议用实际形状来推一遍。假设输入的 token 序列是(batch_size32, block_size128, d_model256)。首先输入会经过三个独立的线性层生成 Q、K、V 三个向量。每个线性层的权重形状都是(256, 256)所以 Q、K、V 的形状和输入一样都是(32, 128, 256)。在多头注意力中这 256 维会被切成 4 个头每个头 64 维也就是把最后一维拆成(4, 64)。Attention 的核心公式是Attention(Q, K, V) softmax(Q K^T / sqrt(d_head)) VQ 和 K 的转置做点积后得到的是一个(k, k)的注意力分数矩阵其中k是序列长度 128。这个矩阵第 i 行第 j 列的含义是在生成第 i 个位置的输出时应该分配给第 j 个位置的“注意力权重”。除以sqrt(d_head)是为了防止点积结果过大导致 softmax 进入饱和区、梯度消失。d_head64所以除以 8。这里必须实现因果掩码causal mask也就是让第 i 个位置只能看到它自己和它之前的位置不能看到未来信息。实现方式是构造一个上三角矩阵把未来位置对应的分数设为一个很大的负数比如-inf这样 softmax 之后这些位置的权重就是 0。def causal_attention(q, k, v, maskNone): d_head q.shape[-1] scores q k.transpose(0, 1, 3, 2) / mx.sqrt(mx.array(d_head)) if mask is not None: scores scores mask weights mx.softmax(scores, axis-1) return weights vMLX 的矩阵运算和 NumPy 很像所以这里的手写实现很直白。如果你追求性能MLX 也提供了nn.MultiheadAttention封装内部实现会更高效但手写一遍能帮你真正理解维度变化。3.3 前馈网络与残差连接每个 Transformer block 还包括一个前馈网络FFN。FFN 由两个线性层组成中间夹一个激活函数结构是d_model - d_ff - d_model。我用了 GELU 激活函数而不是传统的 ReLU因为 GELU 在输出接近 0 的地方更平滑对训练稳定性友好一些。d_ff 我设为 480比常用的 4 倍 d_model 略小。残差连接和 LayerNorm 的排布方式我选了 Pre-LN 结构也就是先做 LayerNorm再进入 Attention/FFN 子层最后加上残差。用公式表达是x x attention(ln1(x)) x x ffn(ln2(x))Pre-LN 相比 Post-LN 在训练时更稳定尤其对学习率不那么敏感这是 GPT 系列实践验证过的经验。每层还加了 dropout0.1防止小模型在小数据上过拟合太快。我把 6 个这样的 block 串起来最后跟一个 LayerNorm然后接一个输出层。输出层用的是Linear(256, 512)没有额外 bias输出直接作为下一个 token 的 logits。为了让参数更紧凑我用了一个常见的技巧输出层复用 Token Embedding 的权重也就是 weight tying。这样模型就省掉了整整 13 万个参数计算量也更小。3.4 参数总量核算321 万是怎么来的现在可以把每一块的参数数量列一个表方便对照。下表基于词表大小 512、d_model256、d_ff480、6 层、序列长度 128模块计算公式参数量Token Embedding512 × 256131,072位置编码128 × 25632,768每层 Attention4 × 256 × 256Q/K/V/O262,144每层 FFN256×480 480×256245,760每层 LayerNorm ×1Pre-LN 两个子层各一个每个 2×2562 × (2 × 256)512每个 Block 合计以上相加508,4166 个 Block 合计508,416 × 63,050,496最终 LayerNorm2 × 256512输出层权重共享weight tying0总计3,214,848这里有个细节要说清楚我在每层 Block 里的 LayerNorm 参数量其实是2 × (2 × 256) 1,024上一行表格里写的是“两个子层各一个”对刚算了 512让我重新核对一下。Pre-LN 中 attention 子层前有一个 LayerNormFFN 子层前也有一个 LayerNorm所以每个 Block 有两个 LayerNorm每个 LayerNorm 有 weight 和 bias 各 256 个参数合计 1,024。上一张表格里写 512 是笔误实际应该是每个 Block 包含两个 LayerNorm参数量为 1,024。修正后的核算模块计算公式参数量Token Embedding512 × 256131,072位置编码128 × 25632,768每层 Attention4 × 256 × 256262,144每层 FFN256×480 480×256245,760每层 LayerNorm ×22 × (2×256)1,024每个 Block 合计508,9286 个 Block 合计508,928 × 63,053,568最终 LayerNorm2×256512总计3,217,920所以最终参数量是 3,217,920也就是约 321 万。这个数字我还是比较满意的它小到可以快速训练大到能展现一个完整 Transformer 所拥有的全部结构细节。你如果改动了 d_model、层数、词表大小任何一个值都可以用同样的公式重新算一遍。4. 训练循环与关键超参调优4.1 损失函数交叉熵训练语言模型本质上是一个分类问题在每个位置上从 512 个可能的 token 中预测出正确的下一个 token。损失函数用的是交叉熵Cross Entropy它衡量的是模型预测的概率分布和真实分布之间的差距。在 MLX 里交叉熵的写法很直接。模型 forward 得到 logits 后形状是(batch_size, block_size, vocab_size)reshape 成(batch_size * block_size, vocab_size)然后和 targetreshape 成同样长度的整数张量一起传给mx.losses.cross_entropy。def compute_loss(model, x, y, dropoutTrue): logits, _ model(x, dropoutdropout) B, T, V logits.shape logits logits.reshape(-1, V) targets y.reshape(-1) loss mx.losses.cross_entropy(logits, targets, reductionmean) return loss这里有个小细节reductionmean是求所有位置的平均 loss不是求和。用平均的好处是即使你改了 batch_size 或者 block_sizeloss 的量级不会变方便对比不同配置的训练效果。我没有用 label smoothing。虽然 label smoothing 可以让模型不要过度自信但对这种小模型小数据集效果并不明显反而会让困惑度稍微变差。为了演示清晰直接用原始交叉熵。4.2 优化器与学习率调度优化器我选了 AdamW这是目前训练 Transformer 最常用的选择。Adam 的优点是自适应学习率对大多数任务都不需要精细调参AdamW 在此基础上把权重衰减和梯度更新解耦能更有效地抑制过拟合。我的超参配置如下基础学习率3e-4beta1 / beta20.9 / 0.999weight_decay0.1warmup 步数200学习率调度cosine 衰减从 3e-4 衰减到 3e-5为什么要 warmup因为训练刚开始时梯度方向是很不稳定、噪声很大的如果一上来就用很大的学习率容易把参数一下子推到一个很差的位置后面很难恢复。先用小学习率“热身”几百步等梯度方向稳定了再加大学习率能明显提高训练的稳定性。cosine 衰减的作用是训练后期接近收敛时用一个很小且缓慢下降的学习率让模型在最小值附近精细调整不容易来回震荡。MLX 实现这个调度很简单每一步根据当前步数计算当前学习率并更新到优化器即可。def lr_schedule(step, warmup200, max_lr3e-4, min_lr3e-5, total_steps5000): if step warmup: return max_lr * (step 1) / warmup progress (step - warmup) / max(1, total_steps - warmup) return min_lr 0.5 * (max_lr - min_lr) * (1 mx.cos(mx.pi * progress))4.3 训练循环写法MLX 的训练循环写法和 PyTorch 很接近但更精简。核心是用nn.value_and_grad同时计算损失和梯度然后调用优化器的update方法更新参数。optimizer optim.AdamW(learning_ratelr_schedule(0), betas(0.9, 0.999)) loss_and_grad nn.value_and_grad(model, compute_loss) for step in range(total_steps): xb, yb get_batch(train_data, batch_size32, block_size128) lr lr_schedule(step) optimizer.learning_rate lr loss, grads loss_and_grad(model, xb, yb) optimizer.update(model, grads) mx.eval(model.parameters(), optimizer.state)这里有个容易忽略的点mx.eval。MLX 是惰性求值的optimizer.update只是把计算图构造好真正执行要等遇到mx.eval。如果你忘了这一步训练循环会越跑越慢内存不断累积。所以在每一步末尾一定记得调用mx.eval这是 MLX 和 PyTorch 在写法上最大的区别之一也是新手最容易踩的坑。另外我把 batch_size 设为 32每次喂进去的数据是32 × 128 4096个 token。对 3.21M 参数的小模型来说这个 batch 大小在 MacBook 上跑起来很轻松一步大概几十毫秒到一百多毫秒。我记得当 loss 从初始的 6.2 左右开始下降前 500 步会掉得特别快到 2000 步左右 train loss 能降到 1.5 以下val loss 大概在 1.6-1.7 之间。这个阶段如果去采样生成已经能产出一些看起来有点语法结构的文本了虽然内容毫无逻辑可言。4.4 训练中的观察从 loss 曲线能读出什么很多人训练的时候只盯着总 loss 一个数字这其实远远不够。我会同时记录 train loss 和 val loss并在每 100 步输出一次。这两条曲线能告诉你很多信息如果 train loss 持续下降但 val loss 过了某个点开始回升说明模型在过拟合——它在背训练数据而不是学习通用规律。这时候应该增大 dropout、增加数据量或者提前停止训练。如果两条 loss 都居高不下、下降很慢可能是学习率太小或者数据预处理出了问题。我会检查 token 序列有没有错位、词表是否覆盖了数据里的高频片段。如果 loss 出现剧烈震荡甚至变成 NaN最常见的原因是学习率太大。尤其是训练初期一个过大的梯度更新可能直接把参数推出正常范围导致后续梯度爆炸。我的经验是遇到 NaN先看学习率再查梯度范数。我在训练脚本里加了一个简单的梯度裁剪逻辑如果梯度范数超过 1.0就按比例缩放梯度。这个操作虽然简单但对防止训练后期不稳定非常有效。if grads_norm max_grad_norm: grads {k: v * (max_grad_norm / grads_norm) for k, v in grads.items()}5. 常见问题与排查技巧实录5.1 Loss 不下降或爆成 NaN我这次训练过程中遇到的第一个异常就是 loss 在前 100 步卡在 6.0 附近不动。检查了一遍发现问题出在数据切分我在构造 target 时把整个序列左移了两位等于让模型预测“下一个的下一个”这当然学不出来。修正之后 loss 才开始正常下降。这类问题最坑的地方在于它不会报错表面看起来一切正常只有通过检查样本才能发现。如果你也遇到 loss 不降建议按这个顺序排查打印一条 input 和 target人工检查目标是不是正确的下一个 token。验证初始 loss 是否符合预期。比如词表大小 512随机初始化时交叉熵应该在log(512) ≈ 6.24左右。如果初始 loss 偏离这个值太多说明数据或模型初始化有异常。用一个极小数据集几十条样本试跑几百步看 loss 能否降到接近 0。如果能说明模型和数据链路没问题问题出在优化器或学习率上。引入梯度裁剪排除梯度爆炸的可能。5.2 训练速度慢或内存占用高MLX 在 Mac 上跑小模型一般不会涉及显存不足的问题统一内存的机制让它能灵活使用整台机器的内存。但如果你把 batch_size 调到 256 以上同时 block_size 又设成 256每一步的注意力矩阵就是256 × 256 × 256 × 256这种量级内存占用会迅速上升。我在调试时观察到当 batch_size32、block_size128 时训练过程的峰值内存大概在 2GB 左右如果把 batch_size 翻到 128内存会涨到 5GB 以上。对于 8GB 统一内存的基础款 MacBook这个配置就跑得比较吃力了。调参时我会优先看两个方向一是减小 batch_size二是减小 block_size。损失一点训练吞吐换回稳定性是值得的。另外我发现MLX 的惰性求值如果配合不好会导致内存不断累积。具体表现是训练刚开始内存占用正常越跑越高最后被系统杀掉。解决办法就是前面反复提过的mx.eval一定要在每次 optimizer update 后强制触发实际计算把中间结果释放掉。5.3 训练后期 loss 在震荡怎么办训练到中后段train loss 降到 1.3 左右之后我一度观察到 loss 在小范围内来回抖动不稳定地上下波动。我当时做了三件事效果显著第一把基础学习率从 3e-4 下调到 2e-4并把 cosine 衰减的最低点设为 1e-5让后期更平稳。第二检查 batch 的随机性——如果数据顺序没有充分 shuffle连续几个 batch 内容相近会产生虚假的波动。我确认了每次 get_batch 都是随机采样排除了这个因素。第三增大 batch_size 从 32 到 64让每步梯度估计更准确这一步对抑制震荡帮助最大。5.4 一个问题排查速查表我把这次训练踩过的坑整理成了一个小表以后排查问题会方便很多现象可能原因解决方案loss 不降数据错位、学习率过小、tokenizer 有问题打印样本检查调大学习率重建词表loss 变成 NaN学习率过大、梯度爆炸降低学习率引入梯度裁剪train loss 低但 val loss 高过拟合增大 dropout增加数据提前停止loss 震荡batch 太小、学习率偏高、数据未充分 shuffle增大 batch降低学习率检查采样逻辑训练越来越慢缺少 mx.eval、内存累积在更新后调用 mx.eval生成全是重复内容温度过低、模型欠拟合/过拟合调高 temperature检查训练曲线6. 推理验证模型到底学到了什么6.1 写一个采样生成函数训练完的模型不能只盯着 loss 看真正好玩的是让它“续写”文本。采样生成的逻辑和训练时不一样训练时是一次性输入一整段文本并行预测所有位置生成时只能从左到右逐 token 进行因为当前预测依赖之前生成的所有 token。实现上我先输入一个固定的起始 token比如\n然后循环做 forward取最后一个位置的 logits经过 softmax 变成概率分布再按概率采样得到下一个 token把它拼到序列尾部继续循环。这里一个关键超参是 temperature它控制采样的随机程度。temperature 越低越倾向选概率最高的 token越高越可能选中概率低但更“有创意”的 token。def generate(model, start_tokens, max_new_tokens200, temperature0.8, top_k50): model.eval() tokens list(start_tokens) for _ in range(max_new_tokens): x mx.array([tokens[-block_size:]]) # 只取最后 block_size 个 logits, _ model(x, dropoutFalse) logits logits[0, -1, :] / temperature # top-k 截断 topk_logits mx.topk(logits, top_k) mask logits topk_logits[-1] logits mx.where(mask, mx.array(-float(inf)), logits) probs mx.softmax(logits, axis-1) next_token mx.random.categorical(probs.log()).item() tokens.append(next_token) return tokenstop-k 采样是为了避免每次都从整个词表里采样导致小概率的不相关内容频繁出现。我只保留概率最高的前 50 个 token从这个子集里重新归一化后采样。这是目前生成文本最常用的策略之一和 top-p核采样搭配使用效果更好。小模型我建议 top-k 设 40-60 之间temperature 设 0.7-0.9生成的文本既有变化又不至于完全乱来。6.2 实际生成的样本与分析我训练了大约 5000 步之后用\n开头让模型生成了一段“莎翁风格”的文本。虽然只是玩具模型它已经能模仿出一些语法结构了。我摘一小段出来KING: The duke of Gloster, and the duke of York, Have been in the kingdom by the sea. What would you have me do?从文本片段可以看出模型确实学到了几个层面的东西。词汇层面它学会了一些常见的单词搭配像 “the duke of” 这种高频词组被完整地记住了。语法层面它学会了一些基本的词序规律比如主语后面跟动词。这些规律本质上是 BPE 词表和高频模式的统计结果模型并没有真正理解“公爵”是什么但已经能基于上下文统计规律生成合理结构了。遗憾的地方也很明显。句子之间的逻辑基本不存在人物关系混乱甚至会出现前后矛盾的说法。这说明 321 万参数和 1MB 数据的组合还远不足以让模型建立长程的依赖关系。如果你做同样的训练用不同的随机种子或数据生成的文本风格会非常不一样这个小实验本身就很有观察价值。6.3 从 321 万到更大模型扩展路线训练完这个小模型之后我觉得最值得做的一件事情是把这套流程迁移到更大的模型和数据上。MLX 的好处是同一套代码可以在内存允许的范围内平滑扩大模型规模。你可以先尝试把 d_model 提到 512、层数加到 8看参数总量变成多少再看训练速度变化多大。如果你感兴趣的方向其实是时下流行的 LoRA 微调低秩适配我也建议先把这个小项目跑熟再上。LoRA 的原理是冻结原模型参数只训练一小部分低秩矩阵作为增量很多工具链都已经把这条流程封得很易用网上能找到大量现成方案。但如果没有亲手训练过一个完整模型你对“哪些层参与了学习、哪些层实际上被冻结”不会形成直观感知。等你理解了底层的梯度流动逻辑再去玩 LoRA 或者其他适配方案理解深度完全不一样。另外MLX 也可以直接加载和运行一些开源量化模型做推理比如社区里经常提到的 qwen 系列小尺寸版本。这类模型已经有很成熟的转换脚本和示例你完全可以在本地跑一个更大的模型做对比。见过几百亿参数模型的输出质量再回头看你亲手训的 321 万小模型那种“差的不是一点两点而是几个数量级”的震撼感比任何论文图表都直观。写在最后我最大的体会是训练一个 321 万参数的小模型比读十篇 Transformer 原理文章都管用。原理文章告诉你 Attention 是“加权求和”但不会告诉你 QKV 矩阵的形状在多头切分时怎么 reshape文章告诉你“残差连接能防止梯度消失”但不会告诉你 Pre-LN 和 Post-LN 在实际训练中收敛速度差了多少。这些细节只有亲手敲过代码、盯着 loss 曲线看过几十次之后才能真正变成你自己的知识。最后再分享一个小技巧训练完千万不要只盯着最终 loss一定要多生成几次文本看看效果。loss 是个冷冰冰的标量但生成文本是模型学习成果的直观体现。你会在一次次生成结果里真切看见模型的成长。这也是整个项目中最让人上瘾的部分。
返回列表