从零训练MiniMind:手把手跑通大模型预训练全流程

发布时间:2026/10/12 4:55:20
从零训练MiniMind:手把手跑通大模型预训练全流程
学习笔记写到第十篇终于到了我最喜欢的环节动手把一个小模型从零训出来。前几篇一直在讲注意力机制、位置编码、分词原理知识点铺了一桌子但只看不练总觉得隔了一层。这篇笔记里的 MiniMind 指的是一类把大语言模型缩到极简程度的教学项目参数少则几百万、多则几千万但数据准备、分词、预训练、推理这条完整链路一个都不少。我这一篇的目标很朴素在你自己的笔记本上把这条链路完整跑一遍最后得到一个真正能续写文字、能聊上几句的小模型。适合那种还没亲手训过模型、想搞懂大模型训练全流程的朋友也适合已经跑过一些现成代码、但想弄明白每个环节为什么这么设计的人。1. 训练前先定三件事显存、数据量、验收标准很多人一上来就问“用什么模型”、“怎么调参”其实训练一个小模型之前最该想清楚的是你的硬件底线、数据预算和“什么样算成功”。这三件事定了后面的代码就是往框架里填东西。1.1 显存不是越大越好而是够用就好我手头用的是一块 8GB 显存的普通消费级显卡十几年前可能觉得这是顶配现在跑大模型只能算入门。但训练 MiniMind 恰恰不需要多夸张的设备关键是把参数量和激活值算明白。先给一个硬估算。假设模型配置是 8000 词表、512 维隐藏层、8 层 Transformer、序列长度 512整体参数量大约在 3700 万左右。用 fp32 存权重需要不到 150MB加上 AdamW 优化器的参数副本也才 600MB 上下看起来非常轻松。真正吃显存的是前向传播和反向传播过程中的激活值这个跟 batch size、序列长度、层数直接相关。序列长度 512、batch size 8 的时候激活值通常要占 3GB 到 5GB8GB 显卡刚好在红线附近。我给新手一个不费脑的配置参考表配置档位参数量训练数据量显存要求大致耗时消费级 GPU入门跑通150万200万 token4GB1到2小时标准练习3700万2000万 token8GB半天到一天进阶尝试1亿以上1亿 token16GB数天如果你连独立显卡都没有CPU 也能跑只是时间要按十倍往上翻。用最小的 150 万参数配置跑几百个 step理解流程完全够用。1.2 数据量别照抄缩放定律练习要打折大模型领域有个著名的经验法则叫“20 tokens per parameter”意思是参数量 3700 万的模型理论上需要 7.4 亿 token 的数据。这个数字对个人学习者来说太不现实了光下载、清洗、预处理就可能劝退一半的人。我的建议是把这条法则当上限而不是标准。练习项目的主要目的是理解训练流程和观察模型的渐进变化数据量可以放宽到参数的 5 到 10 倍。3700 万参数配 2000 万到 4000 万 token就是比较合理的练习区间。数据量再少的话loss 也能降但模型只能记住高频短语谈不上学习语言规律。还有一个比数据量更重要的点数据质量。如果语料里有大量重复段落、乱码字符、超长空行模型会把“复读”当成规律来学。清洗时至少要做三件事去重、去广告标记、按句子边界切分。我这次用的就是一份约 25MB 的公开中文语料覆盖新闻、散文和简单问答token 数大概在 2200 万左右跑起来正合适。1.3 验收标准不是“loss 越小越好”训练前我强烈建议你先写下一句话作为验收基准比如“给定开头‘那天傍晚’模型能续写出语法通顺、和天气或心情相关的一句话”。为什么非要这么干因为 loss 只是一个统计量它下降只代表模型在训练集上的预测概率变大不代表它真的理解了语言。我会同时盯三个信号。第一个是训练 loss 是否平滑下降如果出现长期不动或者突然 NaN那是实现有问题。第二个是验证集 loss在训练中每隔一段时间跑一批没见过的数据如果验证 loss 反而上升说明模型开始死记硬背训练集。第三个才是人工验收自己写几个不同风格的开头看看生成结果是否正常。记住这是练习项目不是要跟大模型比能力目标是跑通链路并且知道每个指标在说什么。2. 从文本到样本分词器、特殊标记与样本打包有了原始语料下一步就是让模型“吃”下去。这一步的核心是把字符串变成 token id 序列再切成训练样本。很多新手直接拿字符来训练不是不行但效率低得离谱我劝你别走这条路。2.1 为什么不能直接拿字符喂模型如果不分词把每个汉字当一个 token词表大概是两三千字看似更简单但模型学不到“词语”这种高层概念。比如“机器学习”四个字字符级模型必须自己从零发现“机器”“学习”经常连着出现这会大幅增加训练难度。更重要的是英文单词的变形go、going、gone在字符级下完全没有共享信息。子词分词是现在的标准方案一个单词或者一个常用词根就是一个 token。中文场景下一个词表里既要有常见汉字也要有“人工”“智能”这类高频词还要能通过组合字符来覆盖生僻词。MiniMind 这类小模型没必要直接用大模型的现成词表因为大模型词表动辄五万十万嵌入层就会吃掉大量参数。自己训练一个 8000 到 16000 大小的词表才是小模型的正确姿势。2.2 用 BPE 训练自己的词表BPEByte Pair Encoding的核心思路不复杂从字符开始反复统计相邻两个单元的出现频率每次把最高频的一对合并成新的单元直到词表达到目标大小。这个算法成熟跑起来也快。下面是训练一个 8000 词表的示例代码我用的是 tokenizers 这个开源库from tokenizers import Tokenizer, models, trainers, pre_tokenizers # 初始化一个 BPE 模型 tok Tokenizer(models.BPE()) # 预分词器ByteLevel 可以处理任意 Unicode 字符 tok.pre_tokenizer pre_tokenizers.ByteLevel(add_prefix_spaceFalse) # 指定特殊 token顺序很重要id 从 0 开始 special [pad, unk, s, /s] trainer trainers.BpeTrainer( vocab_size8000, special_tokensspecial, min_frequency2, ) # 训练并保存 tok.train(files[corpus.txt], trainertrainer) tok.save(minimind_tokenizer.json)几个容易踩的细节。vocab_size8000对中文语料是一个比较经济的值模型参数的很大一部分都花在嵌入层上词表越大嵌入层就越重。min_frequency2表示出现少于两次的字符组合不合并可以有效防止词表被生僻字污染。pad、unk、s、/s这四个特殊 token 的顺序固定后就不要改因为模型训练完会记住它们的 id换顺序等于换字典。2.3 把文本切成固定长度的训练样本原始文本长度不一Transformer 训练时一般要求固定序列长度。最粗暴的做法是按 512 个 token 硬切但很容易把一句话拦腰切断模型学到的跨句模式是畸形的。更好的做法是先按标点把文本切成小段再尽量拼满 512 的长度。简单实现长这样def build_samples(text, tokenizer, seq_len512): # 先按句子切分保证每个片段完整 sentences split_sentences(text) # 自定义函数按。切 buffer [] samples [] for sent in sentences: ids tokenizer.encode(sent).ids if len(buffer) len(ids) seq_len - 2: # 当前 buffer 拼成一个样本 sample [tokenizer.token_to_id(s)] buffer sample sample[:seq_len - 1] [tokenizer.token_to_id(/s)] samples.append(sample) buffer [] else: buffer.extend(ids) return samples注意每个样本前后要加s和/s。开头标记告诉模型“一段话从这里开始”结尾标记给模型一个停下来的信号。如果没有结尾标记续写时模型会一直生成文本结束得不明不白。最终一个 batch 的数据张量是三个input_ids形状(B, L)表示每个位置上的 token id。attention_mask形状(B, L)1 表示真实 token0 表示 padding训练时要通过掩码把这些位置排除掉。labels形状(B, L)预测目标。预训练遵循“预测下一个 token”规则labels[i, t] input_ids[i, t1]最后一个位置没有下一个 token通常用 -100 占位计算损失时会自动忽略。这套设计是整个自监督训练的基石。模型做的事情本质上是猜谜看到前 511 个 token预测第 512 个是什么然后是看到前 512 个预测第 513 个直到跑完整个序列。3. MiniMind 的骨架几行代码搭出一个小 Transformer数据准备好了接下来是模型本体。我不建议直接抄一个 Transformer 库的完整实现然后黑盒调用自己动手写一遍前向传播你对每一层的作用会有完全不同的理解。3.1 整体配置我把 MiniMind 的配置定义成一个数据类方便保存和加载from dataclasses import dataclass dataclass class MiniMindConfig: vocab_size: int 8000 # 词表大小 dim: int 512 # 模型隐藏层维度 n_layers: int 8 # Transformer 层数 n_heads: int 8 # 注意力头数 head_dim: int 64 # 每个头的维度 ff_hidden: int 2048 # 前馈层隐藏维度 seq_len: int 512 # 序列长度 dropout: float 0.0 # 小模型一般不用 dropout为什么选 512 维、8 层、8 头这个组合是“能学到东西又不至于跑不动”的甜点区。维度太低表达能力不足训练半天 loss 下不去维度太高激活和参数同时膨胀消费级显卡直接爆炸。层数同理8 层对于几千万参数的模型已经足够学习复杂的语言层次。头数这里等于 8每个头的维度是 64如果以后想换成 Grouped Query Attention把每个头的维度调大、把共享的头数减少就行。3.2 Embedding、输出层和旋转位置编码模型输入是一串 token id第一步把它映射成稠密向量。这步就是查表权重形状是(vocab_size, dim)。一个常用的技巧是让输入嵌入层和输出 logits 层共享同一份权重因为输入输出本质上都在同一个 token 空间里共享能减少大量参数。MiniMind 的位置编码用的是 RoPE旋转位置编码。它不像传统位置编码那样把位置信息加在向量上而是对每个位置的向量做旋转操作旋转角度和位置相关。这样设计的好处是不同位置之间的相对位移天然体现在向量的夹角差上模型更容易学习相对位置关系。import torch import math def precompute_rope(seq_len, head_dim, base10000.0): # 计算每个位置的角度 inv_freq 1.0 / (base ** (torch.arange(0, head_dim, 2).float() / head_dim)) t torch.arange(seq_len).float() freqs torch.outer(t, inv_freq) # (seq_len, head_dim // 2) return freqs def apply_rope(x, freqs): # x 形状: (B, heads, L, head_dim) x1 x[..., : freqs.shape[-1]] x2 x[..., freqs.shape[-1]:] cos freqs.cos() sin freqs.sin() rot_x1 x1 * cos - x2 * sin rot_x2 x1 * sin x2 * cos return torch.cat([rot_x1, rot_x2], dim-1)这里我简化成只旋转一半维度完整实现会把旋转后的结果按更细致的方式交错合并。新手不需要抠细节只要理解“不同位置有不同的角度位置的差值决定了两组向量之间的旋转角差”就够用了。3.3 注意力模块和前馈模块核心的 Self-Attention 层我用一个线性层同时生成 Q、K、V然后再拆分这样代码更紧凑计算也更高效class Attention(nn.Module): def __init__(self, dim, n_heads, head_dim): super().__init__() self.n_heads n_heads self.head_dim head_dim self.qkv nn.Linear(dim, 3 * n_heads * head_dim, biasFalse) self.wo nn.Linear(n_heads * head_dim, dim, biasFalse) def forward(self, x, freqs, causal_mask): B, L, _ x.shape qkv self.qkv(x) # (B, L, 3 * n_heads * head_dim) q, k, v qkv.chunk(3, dim-1) q q.view(B, L, self.n_heads, self.head_dim).transpose(1, 2) k k.view(B, L, self.n_heads, self.head_dim).transpose(1, 2) v v.view(B, L, self.n_heads, self.head_dim).transpose(1, 2) q apply_rope(q, freqs) k apply_rope(k, freqs) scores q k.transpose(-1, -2) / math.sqrt(self.head_dim) scores scores.masked_fill(causal_mask 0, float(-inf)) attn torch.softmax(scores, dim-1) out attn v out out.transpose(1, 2).reshape(B, L, self.n_heads * self.head_dim) return self.wo(out)因果掩码causal_mask是一个上三角为 0 的矩阵作用是让当前位置只能看到它自己和它之前的位置。比如第 5 个 token 不能看到第 6 个 token 的信息不然就相当于把答案泄露给模型了。这个掩码在训练时是必需的推理时因为有 KV Cache只需要生成当前最后一个位置掩码就不那么关键。为什么要除以sqrt(head_dim)因为两个维度为 64 的向量做点积数值随维度增大而变大softmax 对大数值非常敏感稍微变大一点就趋近于 0 和 1梯度会变得很小甚至消失。除以根号 64 等于 8把得分拉回比较平缓的范围。前馈层我用了 SwiGLU 结构它比经典的两层线性加 ReLU 多了一个门控分支。公式是out W2(swish(W1(x)) * W3(x))swish就是x * sigmoid(x)是 ReLU 的平滑版本。多出的W3相当于一个“开关”控制W1的输出哪些要保留、哪些要抑制。虽然多三分之一的参数但对小模型来说效果提升是值得的。class FeedForward(nn.Module): def __init__(self, dim, ff_hidden): super().__init__() self.w1 nn.Linear(dim, ff_hidden, biasFalse) self.w2 nn.Linear(ff_hidden, dim, biasFalse) self.w3 nn.Linear(dim, ff_hidden, biasFalse) def forward(self, x): return self.w2(torch.nn.functional.silu(self.w1(x)) * self.w3(x))每一层 Transformer 的结构都是“注意力前馈层”外面再用 RMSNorm 归一化最后加残差连接。残差保证了梯度可以跨层流动RMSNorm 则把激活值拉到稳定的尺度这两个组件少了任何一个模型都很难训练。3.4 参数量估算和显存预算有了结构可以算一下的参数量输入输出共享的嵌入层8000 * 512 410万每层 AttentionQKV 矩阵 512 * 1536 78.6万输出矩阵 512 * 512 26.2万合计约 105万每层 SwiGLU 前馈层512 * 2048 * 3 315万每层合计约 420万8 层就是 3360万加上嵌入层总数约 3700万用 fp32 训练模型权重 148MBAdamW 的动量、方差等状态约为权重量的 3 倍共约 600MB。但激活值会随着 batch size 和序列长度增长实测下来 batch size 8、序列长度 512 时总显存大概在 5GB 左右8GB 显卡能跑但很局促。如果新手只想跑通我建议把模型缩到 2 层参数量降到 1500 万以内训练速度会快很多。4. 训练循环里的隐形开关损失计算、学习率调度与梯度裁剪模型定义好只是骨架训练循环才是真正决定模型能不能学会的关键。这个环节里几个隐形开关几乎决定了训练的成败。4.1 损失函数预测下一个 token预训练的损失函数永远只有一个目标让模型对正确的下一个 token 输出更高概率。前向传播后得到形状为(B, L, vocab_size)的 logits把第 0 到第 L-1 个位置的预测和第 1 到第 L 个位置的真实 token 对齐就是标准的交叉熵。logits model(input_ids) # (B, L, vocab_size) shift_logits logits[:, :-1, :].contiguous() shift_labels input_ids[:, 1:].contiguous() loss torch.nn.functional.cross_entropy( shift_logits.view(-1, vocab_size), shift_labels.view(-1) )view(-1, vocab_size)这一步把(B, L, vocab_size)展平成二维矩阵相当于把每个位置当作一个独立分类问题来处理PyTorch 会自动对 batch 和序列两个维度求平均。如果某些位置的标签是 -100交叉熵函数会自动忽略这是处理填充位置的常用手法。初始时刻的 loss 大概是多少词表 8000随机初始化模型对所有 token 几乎等概率交叉熵大约是ln(8000)约等于 8.99。如果你看到 loss 一直停在 8.99 附近不动说明模型完全没有学到东西可能是数据喂错了或者梯度没传起来。如果 loss 从 8.99 慢慢降到 4 到 5相当于困惑度从 8000 降到了 55 左右说明高频词已经能稳定预测了。再往下到 3 到 4模型已经掌握了不少语法规律算是一个基本能用的 MiniMind。4.2 优化器、参数分组和 warmup优化器用 AdamW而不用普通 Adam。AdamW 把权重衰减从梯度更新中分离出来相当于显式地“每隔一段时间把权重往零拉一点”防止个别参数无限膨胀。但权重衰减不能应用到所有参数上比如 RMSNorm 的缩放参数和 bias 系数本来就该灵活变化对它们做衰减会限制模型表达能力。所以标准做法是把参数分成两组decay_params [p for p in model.parameters() if p.dim() 1] no_decay_params [p for p in model.parameters() if p.dim() 1] optimizer torch.optim.AdamW([ {params: decay_params, weight_decay: 0.1}, {params: no_decay_params, weight_decay: 0.0}, ], lr3e-4)学习率调度是另一个关键。小模型虽然小但同样需要 warmup。我见过不少新手一上来就全速跑训练到几百步 loss 直接变成 NaN。原因很简单Adam 优化器在最初几步的二阶动量估计非常不准梯度更新容易冲过头。warmup 就是让学习率从 0 线性升到峰值给动量估计一段预热时间。峰值学习率乘上缩放比例我的习惯是if step warmup_steps: lr peak_lr * (step 1) / warmup_steps else: progress (step - warmup_steps) / (total_steps - warmup_steps) lr peak_lr * 0.5 * (1 math.cos(math.pi * progress))这个余弦退火的思路是前半段学习率保持高位让模型快速逼近后半段逐步降低让模型在损失曲面里落进一个比较平缓的坑里。实际训练中4000 步的总步数warmup 用 500 步峰值学习率 3e-4效果比较稳定。4.3 梯度裁剪和断点保存再提一个容易忽略的细节梯度范数裁剪。语言模型训练过程中偶尔会遇到某些样本产生特别大的梯度如果不加约束参数一步就可能跳出正常区域然后 loss 变成 NaN。一行代码就能解决torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)这个操作把整个模型的梯度向量除以一个缩放系数保证所有梯度的总范数不超过 1.0。注意它是在loss.backward()之后、optimizer.step()之前调用的。保存 checkpoint 时我建议把模型参数、优化器状态、配置、步数、学习率一起存成一个字典if step % 1000 0: torch.save({ model: model.state_dict(), optimizer: optimizer.state_dict(), step: step, config: config, tokenizer: tokenizer, }, fckpt_{step}.pt)很多初学者只存model.state_dict()续训时模型能加载但优化器的动量和学习率状态全丢了等于重新热启动训练曲线会突然变乱。把整个字典存下来才是完整的可续训方案。训练日志里还要记录一个“吞吐量”指标也就是每秒处理的 token 数计算方法很简单batch_size * seq_len / step耗时。这个数字能直接告诉你代码还有没有优化空间如果只有几百 token/s说明模型太小、数据加载太慢或者 CPU 和 GPU 之间的搬运出了问题。5. 让小模型开口说话采样、温度与人工验收训练完成后模型文件躺在磁盘里只是冰冷的权重让它开口说话需要理解采样策略。这是训练之后最容易出效果也最容易出笑话的环节。5.1 为什么不能总是选概率最大的 token模型对下一个 token 会输出一组 logits很多人第一反应是取 argmax也就是概率最大的那个 token。这样生成的结果会非常“板正”但很快就陷入重复循环。原因是语言本身的连续性远大于确定性每次都选最大概率相当于用一条直线去拟合一条弯弯曲曲的路径必然跑偏。更自然的做法是采样根据模型输出的概率分布掷骰子概率高的 token 被选中的次数多概率低的偶尔也会被选中。这种随机性反而让生成结果更流畅。为了让采样效果更好还需要三个调节旋钮temperature、top-k、top-p。temperature 控制分布的“尖锐程度”。温度越低概率集中在少数 token 上生成更保守温度越高分布越平坦生成越天马行空。top-k 的思路是只保留概率最高的 k 个 token其余全部过滤。top-p 更自适应它按照概率从高到低累计直到累加值超过 p然后只在这些 token 里重新归一化。def sample_next(logits, temperature0.8, top_k50, top_p0.95): logits logits / temperature if top_k is not None: k min(top_k, logits.size(-1)) top_k_values, _ torch.topk(logits, k) logits[logits top_k_values[:, :, -1].unsqueeze(-1)] float(-inf) if top_p is not None and top_p 1.0: probs torch.softmax(logits, dim-1) sorted_probs, sort_indices torch.sort(probs, descendingTrue) cumsum torch.cumsum(sorted_probs, dim-1) removed_mask cumsum - sorted_probs top_p removed_mask[:, :, 1:] removed_mask[:, :, :-1].clone() removed_mask[:, :, 0] False logits[sort_indices[removed_mask]] float(-inf) probs torch.softmax(logits, dim-1) return torch.multinomial(probs, 1)实际用下来temperature 0.7 到 0.9、top-k 50、top-p 0.95 是一组很通用的组合中文语料下基本能生成流畅通顺的句子。5.2 完整生成循环有了采样函数生成就非常简单了。给定一个开头把它 token 化后喂给模型循环执行以下步骤计算 logits取出最后一个位置的分布采样出一个新 token把它拼到序列末尾再作为下一次输入。def generate(prompt, tokenizer, model, max_new_tokens100, temperature0.8, top_k50, top_p0.95): model.eval() input_ids tokenizer.encode(prompt).ids input_ids torch.tensor([input_ids], deviceconfig.device) with torch.no_grad(): for _ in range(max_new_tokens): logits model(input_ids)[:, -1, :] next_id sample_next(logits, temperature, top_k, top_p) input_ids torch.cat([input_ids, next_id], dim-1) if next_id.item() tokenizer.token_to_id(/s): break return tokenizer.decode(input_ids[0].tolist(), skip_special_tokensTrue)这套实现是带了“KV Cache”的最简版本每次生成都把完整序列重算一遍效率不高但足够理解原理。如果要提速可以缓存每一层的 K 和 V 矩阵只算新 token 的注意力但这会引入更多状态管理新手先从零开始学更稳妥。5.3 如何人工验收训练效果训练时 loss 降到 4.0 以下就值得停下来生成几个句子看看了。我准备了三个固定测试每次训练完都用它们验收续写“那天傍晚我” 希望看到环境描写或事件展开。问答式的语料里给一个“什么是人工智能”看它能否生成通顺但不一定正确的解释。重复开头“今天天气很好我们一起去” 看是否能出现自然的下文。小模型没有海量知识别指望它知道网络热搜或者时事新闻它学的是语言模式。如果它生成的内容只是高频词语的堆砌说明训练数据太少或模型容量不足如果语法基本通顺、偶尔有点小聪明那恭喜你训练链路已经成功了。6. 踩坑实录loss 不降、显存紧张和生成乱码的完整排查链路从零开始写这套训练系统我踩过的坑比顺利的一次要多得多。下面这些排查思路是按顺序走的每一步都对应一类典型问题。6.1 loss 不降的第一步不是调参而是检查数据如果你的 loss 一直在 8.99 附近纹丝不动先别急着激动调学习率大概率是数据或者标签出了问题。我把排查顺序固定成下面这样先打印一个 batch 的input_ids和labels人工看一眼是否错位。常见错误是labels忘了右移或者右移后最后一个位置没有被 -100 覆盖导致模型被迫预测一个无意义的 token。检查attention_mask是否覆盖了所有真实 token。如果 mask 错把有效位置置 0模型相当于瞪着白板学语文。检查数据加载器有没有 shuffle。如果完全不打乱数据训练初期看到的全是同一个主题的文本loss 会在某个值附近反复震荡。确认 loss 计算是否真的作用在logits的最后一个维度上。cross_entropy对维度很挑剔好多人把(B, L, V)直接传进去求和维度错了loss 曲线看着在下降实际上在乱学。最后才去看学习率。初学者一上来就把学习率顶到 1e-3 以上大概率直接 NaN但如果你用的是 fp16 且没有做动态损失缩放也可能因为下溢出现 loss 不变的情况这种情况建议改成 bf16 或者退回 fp32。我印象最深的一次是代码里把labels的右移写成了input_ids本身模型训练全程都在“预测自己当前位置的 token”loss 降得飞快但生成结果全是乱码。loss 指标有时候会骗人必须靠生成样本来验证。6.2 显存爆炸时调整顺序比瞎减模型尺寸更有效显存不够是最常见的硬件问题但很多人第一反应是直接把模型层数砍半我觉得这不一定是最优解。正确的排查顺序应该从激活值入手先把 batch size 降到 1如果显存依然不够说明问题在序列长度或模型本身。再降序列长度从 512 降到 256 或 128。因为注意力分数矩阵的形状是(L, L)显存占用随序列长度近似平方增长这里的收益最大。还不行的再考虑减层数或隐藏维度。但这种结构性改变会影响模型容量有时候会让整个训练结果失去参考价值。另外一个隐蔽的显存泄漏点推理和训练交替时以前计算的中间张量没有释放。在训练循环外多调用一次torch.cuda.empty_cache()能解决不少症状但它只是清理缓存不是根治。真正的根治是用完的中间变量及时释放少保留不必要的计算图。小技巧是在不更新梯度的推理阶段包一层torch.no_grad()反向传播不会为这些计算保存中间节点显存立刻降下一大截。6.3 梯度累积解决 batch 大小不够的问题显存小又想要更大的有效 batch梯度累积是标准答案。具体做法是把一个大 batch 拆成几个小 batch分别前向和反向但暂时不更新参数等累计了几个小 batch 的梯度后再统一更新一次。accum_steps 4 loss loss / accum_steps loss.backward() if (step 1) % accum_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() optimizer.zero_grad()注意两个细节每个小 batch 的 loss 要除以累积步数否则等于把有效学习率放大了 accum_steps 倍学习率调度器的step()也要放在参数更新之后而不是每个小 batch 都更新。梯度累积的本质是用时间换空间速度会变慢但能让你用 4GB 显存跑出 16GB 显卡的效果练习场景下非常实用。6.4 生成乱码和无限重复的怪问题训练结束生成测试时最打击人的是两种情况全是乱码、不断重复。乱码大概率不是模型问题而是分词器编码解码错位。你训练时用的是自己的 BPE 词表推理时如果忘记加载同一个 tokenizer或者特殊 token 顺序对不上解码出来的自然全是火星文。排查方法很简单把一句话编码再解码看是否原样返回。如果这一步都对再看模型输入 id 是否经过了正确的 padding 和截断。无限重复则有三个嫌疑温度过低、数据总量太少、序列长度不足。温度低于 0.5 会让采样几乎变成 argmax重复几乎是必然的。数据太少会导致模型只能记忆高频短语生成时在这些短语之间跳来跳去。序列长度不足则是最难根治的因为模型在训练时从来没有见过超过 512 个 token 的上下文测试时硬让它生成 500 个字到后面它已经“忘了”开头在说什么只能基于最近几句话循环。遇到重复先调温度到 0.8 左右顺手把 top-p 从 0.9 降到 0.8通常能缓解。想彻底解决只能加大数据量或增加序列长度重新训练。写在最后的经验这篇笔记写到这里MiniMind 从零训练的基础闭环就完整了。我自己的最大体会是训练小模型这件事贵在“亲手跑完一遍”。你写数据管线时踩过的坑比看十篇原理文章都管用你亲眼看到 loss 从 8.99 一路降到 4.5比任何教程都能说明问题。如果接下来你想继续往前走我的建议顺序是先把语料换成自己熟悉的领域比如你所在行业的文档、你常看的博客文章重新训练一次这样你能更敏感地判断模型是“学会了”还是“背住了”然后试着做一遍 SFT把预训练得到的模型在几十万条问答对上做监督微调体验一下从“续写器”变成“对话者”的过程最后再做一个简单的评测集固定十句话每次训练完都跑一遍记录 loss 的变化。这三步走下来你对大模型训练的理解会比多数只玩过 API 的人扎实得多。