☰
从零训练MiniMind:手把手跑通大模型预训练全流程
2026/10/12 4:55:13 网站建设 项目流程

学习笔记写到第十篇,终于到了我最喜欢的环节:动手把一个小模型从零训出来。前几篇一直在讲注意力机制、位置编码、分词原理,知识点铺了一桌子,但只看不练,总觉得隔了一层。这篇笔记里的 MiniMind 指的是一类把大语言模型缩到极简程度的教学项目,参数少则几百万、多则几千万,但数据准备、分词、预训练、推理这条完整链路一个都不少。我这一篇的目标很朴素:在你自己的笔记本上,把这条链路完整跑一遍,最后得到一个真正能续写文字、能聊上几句的小模型。适合那种还没亲手训过模型、想搞懂大模型训练全流程的朋友,也适合已经跑过一些现成代码、但想弄明白每个环节为什么这么设计的人。

1. 训练前先定三件事:显存、数据量、验收标准

很多人一上来就问“用什么模型”、“怎么调参”,其实训练一个小模型之前,最该想清楚的是你的硬件底线、数据预算和“什么样算成功”。这三件事定了,后面的代码就是往框架里填东西。

1.1 显存不是越大越好,而是够用就好

我手头用的是一块 8GB 显存的普通消费级显卡,十几年前可能觉得这是顶配,现在跑大模型只能算入门。但训练 MiniMind 恰恰不需要多夸张的设备,关键是把参数量和激活值算明白。

先给一个硬估算。假设模型配置是 8000 词表、512 维隐藏层、8 层 Transformer、序列长度 512,整体参数量大约在 3700 万左右。用 fp32 存权重需要不到 150MB,加上 AdamW 优化器的参数副本也才 600MB 上下,看起来非常轻松。真正吃显存的是前向传播和反向传播过程中的激活值,这个跟 batch size、序列长度、层数直接相关。序列长度 512、batch size 8 的时候,激活值通常要占 3GB 到 5GB,8GB 显卡刚好在红线附近。

我给新手一个不费脑的配置参考表:

配置档位参数量训练数据量显存要求大致耗时(消费级 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 训练自己的词表

BPE(Byte 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_space=False) # 指定特殊 token,顺序很重要,id 从 0 开始 special = ["<pad>", "<unk>", "<s>", "</s>"] trainer = trainers.BpeTrainer( vocab_size=8000, special_tokens=special, min_frequency=2, ) # 训练并保存 tok.train(files=["corpus.txt"], trainer=trainer) tok.save("minimind_tokenizer.json")

几个容易踩的细节。vocab_size=8000对中文语料是一个比较经济的值,模型参数的很大一部分都花在嵌入层上,词表越大,嵌入层就越重。min_frequency=2表示出现少于两次的字符组合不合并,可以有效防止词表被生僻字污染。<pad>、<unk>、<s>、</s>这四个特殊 token 的顺序固定后就不要改,因为模型训练完会记住它们的 id,换顺序等于换字典。

2.3 把文本切成固定长度的训练样本

原始文本长度不一,Transformer 训练时一般要求固定序列长度。最粗暴的做法是按 512 个 token 硬切,但很容易把一句话拦腰切断,模型学到的跨句模式是畸形的。

更好的做法是先按标点把文本切成小段,再尽量拼满 512 的长度。简单实现长这样:

def build_samples(text, tokenizer, seq_len=512): # 先按句子切分,保证每个片段完整 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 表示真实 token,0 表示 padding,训练时要通过掩码把这些位置排除掉。
  • labels:形状(B, L),预测目标。预训练遵循“预测下一个 token”规则,labels[i, t] = input_ids[i, t+1],最后一个位置没有下一个 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, base=10000.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, bias=False) self.wo = nn.Linear(n_heads * head_dim, dim, bias=False) 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, bias=False) self.w2 = nn.Linear(ff_hidden, dim, bias=False) self.w3 = nn.Linear(dim, ff_hidden, bias=False) 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万
  • 每层 Attention:QKV 矩阵 512 * 1536 = 78.6万,输出矩阵 512 * 512 = 26.2万,合计约 105万
  • 每层 SwiGLU 前馈层:512 * 2048 * 3 = 315万
  • 每层合计约 420万,8 层就是 3360万
  • 加上嵌入层,总数约 3700万

用 fp32 训练,模型权重 148MB,AdamW 的动量、方差等状态约为权重量的 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}, ], lr=3e-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_norm=1.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, }, f"ckpt_{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, temperature=0.8, top_k=50, top_p=0.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, descending=True) 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_tokens=100, temperature=0.8, top_k=50, top_p=0.95): model.eval() input_ids = tokenizer.encode(prompt).ids input_ids = torch.tensor([input_ids], device=config.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_tokens=True)

这套实现是带了“KV Cache”的最简版本,每次生成都把完整序列重算一遍,效率不高但足够理解原理。如果要提速,可以缓存每一层的 K 和 V 矩阵,只算新 token 的注意力,但这会引入更多状态管理,新手先从零开始学更稳妥。

5.3 如何人工验收训练效果

训练时 loss 降到 4.0 以下,就值得停下来生成几个句子看看了。我准备了三个固定测试,每次训练完都用它们验收:

  • 续写:“那天傍晚,我” 希望看到环境描写或事件展开。
  • 问答式的语料里,给一个“什么是人工智能?”看它能否生成通顺但不一定正确的解释。
  • 重复开头“今天天气很好,我们一起去” 看是否能出现自然的下文。

小模型没有海量知识,别指望它知道网络热搜或者时事新闻,它学的是语言模式。如果它生成的内容只是高频词语的堆砌,说明训练数据太少或模型容量不足;如果语法基本通顺、偶尔有点小聪明,那恭喜你,训练链路已经成功了。

6. 踩坑实录:loss 不降、显存紧张和生成乱码的完整排查链路

从零开始写这套训练系统,我踩过的坑比顺利的一次要多得多。下面这些排查思路是按顺序走的,每一步都对应一类典型问题。

6.1 loss 不降的第一步不是调参,而是检查数据

如果你的 loss 一直在 8.99 附近纹丝不动,先别急着激动调学习率,大概率是数据或者标签出了问题。我把排查顺序固定成下面这样:

  1. 先打印一个 batch 的input_ids和labels,人工看一眼是否错位。常见错误是labels忘了右移,或者右移后最后一个位置没有被 -100 覆盖,导致模型被迫预测一个无意义的 token。
  2. 检查attention_mask是否覆盖了所有真实 token。如果 mask 错把有效位置置 0,模型相当于瞪着白板学语文。
  3. 检查数据加载器有没有 shuffle。如果完全不打乱数据,训练初期看到的全是同一个主题的文本,loss 会在某个值附近反复震荡。
  4. 确认 loss 计算是否真的作用在logits的最后一个维度上。cross_entropy对维度很挑剔,好多人把(B, L, V)直接传进去,求和维度错了,loss 曲线看着在下降,实际上在乱学。
  5. 最后才去看学习率。初学者一上来就把学习率顶到 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 的人扎实得多。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询