2 小时从零炼一个 64M 小模型:MiniMind 源码精读与显存踩坑

目标读者 :想「真的搞懂 LLM 怎么炼成」的开发者,而非只会调 API

预计阅读 :18~25 分钟

仓库jingyaogong/minimind

关键词:MiniMind、预训练、SFT、DPO、RLHF、显存、BPE、Next Token Prediction


写在前面:为什么值得啃 MiniMind?

大模型时代,很多人停留在:

text 复制代码
会调 OpenAI API  ≠  知道模型权重从哪来
会跑 LoRA 微调   ≠  理解预训练 / SFT / 对齐整条链路

jingyaogong/minimind 的定位很直白:从 0 开始,用原生 PyTorch(不依赖 HuggingFace Trainer 黑盒),在消费级 GPU 上炼出一个约 64M 参数的小语言模型

官方 README 给出的「入门标杆」是(以 单卡 RTX 3090 为例):

  • 约 2 小时pretrain_t2t_mini + sft_t2t_mini 各 1 epoch
  • 约 3 元人民币:对应租卡成本量级

注:「2 小时 / 3 块钱」主要指 SFT 阶段在 3090 上跑 1 epoch 的实测;完整 Pretrain→SFT→DPO 会更长,但仍在个人可承受范围。

本文按 数据管线 → 分词 → 预训练 → SFT → DPO 五步拆开源码,并给出 12GB+ 显存配置、常见 OOM 与对齐踩坑


一、全景:一条可运行的「炼模型」流水线

text 复制代码
┌─────────────┐    ┌──────────────┐    ┌─────────────┐    ┌─────────────┐    ┌─────────────┐
│ ① 数据管线   │───►│ ② 分词器 BPE  │───►│ ③ 预训练 PT  │───►│ ④ SFT 微调   │───►│ ⑤ DPO 对齐   │
│ jsonl 清洗  │    │ train_tokenizer│    │ Next Token  │    │ 只学回答部分 │    │ chosen/rej  │
└─────────────┘    └──────────────┘    └─────────────┘    └─────────────┘    └─────────────┘
       │                  │                   │                   │                   │
       ▼                  ▼                   ▼                   ▼                   ▼
 pretrain_*.jsonl    vocab≈6400         pretrain_768.pth    full_sft_768.pth    dpo_768.pth
 sft_t2t_mini.jsonl  chat template      学会「接龙」          学会「对话格式」      学会「偏好」
 dpo.jsonl

权重命名规律hidden_size=768 时):

阶段 默认输出 加载方式 --from_weight
Pretrain pretrain_768.pth none(从头)
SFT full_sft_768.pth pretrain
DPO dpo_768.pth full_sft

#mermaid-svg-mUkzOxHeD5QkCJNu{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-mUkzOxHeD5QkCJNu .error-icon{fill:#552222;}#mermaid-svg-mUkzOxHeD5QkCJNu .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-mUkzOxHeD5QkCJNu .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-mUkzOxHeD5QkCJNu .marker{fill:#333333;stroke:#333333;}#mermaid-svg-mUkzOxHeD5QkCJNu .marker.cross{stroke:#333333;}#mermaid-svg-mUkzOxHeD5QkCJNu svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-mUkzOxHeD5QkCJNu p{margin:0;}#mermaid-svg-mUkzOxHeD5QkCJNu .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu .cluster-label text{fill:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu .cluster-label span{color:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu .cluster-label span p{background-color:transparent;}#mermaid-svg-mUkzOxHeD5QkCJNu .label text,#mermaid-svg-mUkzOxHeD5QkCJNu span{fill:#333;color:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu .node rect,#mermaid-svg-mUkzOxHeD5QkCJNu .node circle,#mermaid-svg-mUkzOxHeD5QkCJNu .node ellipse,#mermaid-svg-mUkzOxHeD5QkCJNu .node polygon,#mermaid-svg-mUkzOxHeD5QkCJNu .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-mUkzOxHeD5QkCJNu .rough-node .label text,#mermaid-svg-mUkzOxHeD5QkCJNu .node .label text,#mermaid-svg-mUkzOxHeD5QkCJNu .image-shape .label,#mermaid-svg-mUkzOxHeD5QkCJNu .icon-shape .label{text-anchor:middle;}#mermaid-svg-mUkzOxHeD5QkCJNu .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-mUkzOxHeD5QkCJNu .rough-node .label,#mermaid-svg-mUkzOxHeD5QkCJNu .node .label,#mermaid-svg-mUkzOxHeD5QkCJNu .image-shape .label,#mermaid-svg-mUkzOxHeD5QkCJNu .icon-shape .label{text-align:center;}#mermaid-svg-mUkzOxHeD5QkCJNu .node.clickable{cursor:pointer;}#mermaid-svg-mUkzOxHeD5QkCJNu .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-mUkzOxHeD5QkCJNu .arrowheadPath{fill:#333333;}#mermaid-svg-mUkzOxHeD5QkCJNu .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-mUkzOxHeD5QkCJNu .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-mUkzOxHeD5QkCJNu .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-mUkzOxHeD5QkCJNu .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-mUkzOxHeD5QkCJNu .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-mUkzOxHeD5QkCJNu .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-mUkzOxHeD5QkCJNu .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-mUkzOxHeD5QkCJNu .cluster text{fill:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu .cluster span{color:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-mUkzOxHeD5QkCJNu .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-mUkzOxHeD5QkCJNu rect.text{fill:none;stroke-width:0;}#mermaid-svg-mUkzOxHeD5QkCJNu .icon-shape,#mermaid-svg-mUkzOxHeD5QkCJNu .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-mUkzOxHeD5QkCJNu .icon-shape p,#mermaid-svg-mUkzOxHeD5QkCJNu .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-mUkzOxHeD5QkCJNu .icon-shape .label rect,#mermaid-svg-mUkzOxHeD5QkCJNu .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-mUkzOxHeD5QkCJNu .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-mUkzOxHeD5QkCJNu .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-mUkzOxHeD5QkCJNu :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} out/
trainer/
数据层
pretrain_t2t_mini.jsonl
sft_t2t_mini.jsonl
dpo.jsonl
train_pretrain.py
train_full_sft.py
train_dpo.py
pretrain_768.pth
full_sft_768.pth
dpo_768.pth


二、模型长什么样?64M 从哪来?

minimind-3 主线配置(对齐 Qwen3 生态的 Dense 小模型):

参数 含义
d_model (hidden_size) 768 词嵌入与隐藏维度
n_layers 8 Transformer 层数(偏「深而窄」)
q_heads / kv_heads 8 / 4 GQA 分组查询注意力
vocab_size 6400 小词表,省参数
max_position_embeddings 32768 训练可截断;推理可用 YaRN 外推
总参数量 ≈64M Dense;MoE 版约 198M-A64M

架构示意(单层 MiniMindBlock):

text 复制代码
                    input hidden_states
                            │
                    ┌───────▼───────┐
                    │  RMSNorm      │
                    └───────┬───────┘
                            │
              ┌─────────────▼─────────────┐
              │  Multi-Head Attention     │
              │  + RoPE + FlashAttn(可选)  │
              └─────────────┬─────────────┘
                            │ + residual
                    ┌───────▼───────┐
                    │  RMSNorm      │
                    └───────┬───────┘
                            │
              ┌─────────────▼─────────────┐
              │  SwiGLU FFN (或 MoE)        │
              └─────────────┬─────────────┘
                            │ + residual
                            ▼
                    output hidden_states

核心代码在 model/model_minimind.pyMiniMindForCausalLM 用标准 Causal LM + CrossEntropy 算 loss;MoE 时额外加 router_aux_loss

python 复制代码
# model_minimind.py --- 损失计算核心(节选)
if labels is not None:
    x, y = logits[..., :-1, :].contiguous(), labels[..., 1:].contiguous()
    loss = F.cross_entropy(x.view(-1, x.size(-1)), y.view(-1), ignore_index=-100)

ignore_index=-100 贯穿 Pretrain / SFT / DPO:padding 与不应学习的位置不参与 loss


三、Step 0~1:数据管线与分词器

3.1 数据从哪来?

快速复现(Zero 路线)推荐组合:

文件 用途 建议 max_seq_len
pretrain_t2t_mini.jsonl 轻量预训练 ≈768
sft_t2t_mini.jsonl 轻量 SFT(含少量 Tool Call) ≈768
dpo.jsonl 偏好对齐 默认脚本 1024

Pretrain 样本格式PretrainDataset):

json 复制代码
{"text": "秦始皇是中国历史上第一位皇帝......"}

SFT 样本格式SFTDataset):

json 复制代码
{
  "conversations": [
    {"role": "user", "content": "介绍一下杭州美食"},
    {"role": "assistant", "content": "杭州有西湖醋鱼、片儿川......"}
  ]
}

DPO 样本格式DPODataset):

json 复制代码
{
  "chosen": [
    {"role": "user", "content": "如何礼貌拒绝加班?"},
    {"role": "assistant", "content": "可以这样表达:感谢信任,但今晚已有安排......"}
  ],
  "rejected": [
    {"role": "user", "content": "如何礼貌拒绝加班?"},
    {"role": "assistant", "content": "不行,别烦我。"}
  ]
}

max_seq_lentoken 长度 ,不是字符数。中文约 1 token ≈ 1.5~1.7 字符(README 口径)。

3.2 分词器:BPE + Chat Template

trainer/train_tokenizer.py 从零训练 BPE ,词表约 6400,并支持:

  • <tool_call> / <tool_response>
  • `` 思考标签(与 Qwen 风格对齐)
  • apply_chat_template 多轮对话模板

为什么要自己训 tokenizer?

小模型 + 小词表 = 更少 embedding 参数,64M 预算才能留给 Transformer 层。这也是「从 0 炼」的完整含义之一。

3.3 数据管线源码精读

dataset/lm_dataset.pyPretrain 逻辑极简:

python 复制代码
# Pretrain:整段 text 做 next-token prediction,padding 标 -100
tokens = tokenizer(str(sample['text']), ..., max_length=self.max_length - 2, truncation=True).input_ids
tokens = [bos] + tokens + [eos]
labels = input_ids.clone()
labels[input_ids == pad_token_id] = -100

SFT 的关键 :只对 assistant 回复段 计算 loss(指令段 mask 掉):

python 复制代码
# generate_labels:找到 assistant 段起止,仅该区间 labels ≠ -100
self.bos_id = tokenizer(f'{tokenizer.bos_token}assistant\n', ...).input_ids
self.eos_id = tokenizer(f'{tokenizer.eos_token}\n', ...).input_ids
# 扫描 input_ids,在 bos_id~eos_id 之间写入真实 label

这一步决定了:SFT 不是让模型背用户问题,而是学「在给定上下文下如何回答」


四、Step 2:预训练(Pretrain)------学会「词语接龙」

4.1 在学什么?

无监督 Next Token Prediction :读大量文本,预测下一个 token。

目标:压缩世界知识与语言统计规律进 64M 参数。

4.2 训练脚本结构(train_pretrain.py

text 复制代码
1. init_distributed_mode()     # 可选多卡
2. MiniMindConfig(768, 8层)
3. init_model(from_weight='none')
4. PretrainDataset + AdamW + bf16 autocast
5. loss = logits_loss + aux_loss(MoE时)
6. 梯度累积 accumulation_steps=8
7. 保存 pretrain_768.pth + checkpoints 续训

默认超参(节选)

参数 默认值 说明
batch_size 32 每步样本数
accumulation_steps 8 有效 batch = 32×8 = 256
learning_rate 5e-4 cosine 调度
max_seq_len 340(可改 768) mini 数据建议 768
dtype bfloat16 3090 友好

启动命令

bash 复制代码
git clone https://github.com/jingyaogong/minimind.git
cd minimind
pip install -r requirements.txt
# 下载 dataset 到 ./dataset(见 README / HuggingFace minimind_dataset)

cd trainer
python train_pretrain.py \
  --data_path ../dataset/pretrain_t2t_mini.jsonl \
  --max_seq_len 768 \
  --epochs 1 \
  --batch_size 16 \
  --accumulation_steps 16

验证

bash 复制代码
python eval_llm.py --weight pretrain
# 期望:能续写常识句,但还不会「聊天格式」

五、Step 3:SFT------学会「对话格式 + 助手行为」

5.1 与 Pretrain 的本质区别

Pretrain SFT
数据 纯文本 text 多轮 conversations
Loss 范围 全序列(除 pad) 仅 assistant 段
学到什么 语言建模 指令跟随、角色、工具标签
输入权重 from_weight=pretrain

官方把大体量 sft_t2t.jsonl 描述为带 mid-training 性质:不仅是格式对齐,也继续灌知识分布。

5.2 源码要点(train_full_sft.py

与 Pretrain 共用同一套训练循环,差异在:

python 复制代码
train_ds = SFTDataset(args.data_path, tokenizer, max_length=args.max_seq_len)
model, tokenizer = init_model(lm_config, args.from_weight, ...)  # 默认加载 pretrain

post_processing_chat 会以一定概率去掉空 `` 块,避免模型学会「假思考」。

启动

bash 复制代码
cd trainer
python train_full_sft.py \
  --from_weight pretrain \
  --data_path ../dataset/sft_t2t_mini.jsonl \
  --max_seq_len 768 \
  --epochs 1 \
  --batch_size 8 \
  --accumulation_steps 8

3090 官方估算sft_t2t_mini 1 epoch ≈ 1.1h / ≈1.43¥


六、Step 4:DPO------把论文里的 RLHF 落到可运行代码

6.1 从 PPO 到 DPO:你真正需要懂的一句

经典 RLHF 要训:Reward Model + Actor + Critic,在线采样,显存和工程成本都高。

DPO(Direct Preference Optimization) 在静态偏好对 (chosen, rejected) 上直接优化:

L D P O = − E log ⁡ σ ( β \[ log ⁡ π θ ( y w ∣ x ) π r e f ( y w ∣ x ) − log ⁡ π θ ( y l ∣ x ) π r e f ( y l ∣ x ) ) ] \mathcal{L}_{DPO} = -\mathbb{E}\left\\log \\sigma\\left(\\beta \\left\[\\log \\frac{\\pi_\\theta(y_w\|x)}{\\pi_{ref}(y_w\|x)} - \\log \\frac{\\pi_\\theta(y_l\|x)}{\\pi_{ref}(y_l\|x)}\\right\right)\right] LDPO=−Elogσ(β\[logπref(yw∣x)πθ(yw∣x)−logπref(yl∣x)πθ(yl∣x))]

白话:让策略模型相对冻结的 ref 模型,更偏爱 chosen、更不偏爱 rejectedβ 控制偏离 ref 的力度。

6.2 MiniMind 的 PyTorch 原生实现(精读 train_dpo.py

① 双模型:policy 可训,ref 冻结

python 复制代码
model, tokenizer = init_model(lm_config, args.from_weight, device=args.device)  # full_sft
ref_model, _ = init_model(lm_config, args.from_weight, device=args.device)
ref_model.eval()
ref_model.requires_grad_(False)

② 把一个 batch 拼成 chosen + rejected 前半后半

python 复制代码
x = torch.cat([x_chosen, x_rejected], dim=0)
y = torch.cat([y_chosen, y_rejected], dim=0)
mask = torch.cat([mask_chosen, mask_rejected], dim=0)

③ 算 token 级 log prob,再按 mask 求序列和

python 复制代码
def logits_to_log_probs(logits, labels):
    log_probs = F.log_softmax(logits, dim=2)
    return torch.gather(log_probs, dim=2, index=labels.unsqueeze(2)).squeeze(-1)

def dpo_loss(ref_log_probs, policy_log_probs, mask, beta):
    ref_log_probs = (ref_log_probs * mask).sum(dim=1)
    policy_log_probs = (policy_log_probs * mask).sum(dim=1)
    batch_size = ref_log_probs.shape[0]
    chosen_ref = ref_log_probs[:batch_size // 2]
    reject_ref = ref_log_probs[batch_size // 2:]
    chosen_pol = policy_log_probs[:batch_size // 2]
    reject_pol = policy_log_probs[batch_size // 2:]
    logits = (chosen_pol - reject_pol) - (chosen_ref - reject_ref)
    return (-F.logsigmoid(beta * logits)).mean()

④ 训练时 ref 前向无梯度

python 复制代码
with torch.no_grad():
    ref_outputs = ref_model(x)
    ref_log_probs = logits_to_log_probs(ref_outputs.logits, y)
outputs = model(x)
policy_log_probs = logits_to_log_probs(outputs.logits, y)
loss = dpo_loss(ref_log_probs, policy_log_probs, mask, beta=args.beta)

这就是「论文公式 → 几十行 PyTorch」的完整落地。

6.3 DPO 超参踩坑(官方 README 强调)

参数 默认 建议
learning_rate 4e-8 不要超过 ~5e-8,否则易「遗忘」SFT 能力
beta 0.15 越大越贴 ref,越小越敢偏离
batch_size 4 显存紧再降
from_weight full_sft 必须基于 SFT 权重

启动

bash 复制代码
cd trainer
python train_dpo.py \
  --from_weight full_sft \
  --data_path ../dataset/dpo.jsonl \
  --max_seq_len 512 \
  --batch_size 2 \
  --learning_rate 4e-8 \
  --beta 0.15

6.4 DPO vs PPO(何时用谁)

DPO PPO
数据 静态偏好对,可反复 epoch 在线 rollout,on-policy
显存 policy + ref(约 2× 推理前向) 常 1.5~2× 于 DPO(还有 Critic)
适合 偏好、安全、礼貌 可验证奖励、做题、Agent
MiniMind train_dpo.py train_ppo.py / GRPO / CISPO

README 直言:DPO 对「会不会做题」提升有限,更偏人类价值对齐。
#mermaid-svg-XJ06xylSuyIkYSH4{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-XJ06xylSuyIkYSH4 .error-icon{fill:#552222;}#mermaid-svg-XJ06xylSuyIkYSH4 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-XJ06xylSuyIkYSH4 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-XJ06xylSuyIkYSH4 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-XJ06xylSuyIkYSH4 .marker.cross{stroke:#333333;}#mermaid-svg-XJ06xylSuyIkYSH4 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-XJ06xylSuyIkYSH4 p{margin:0;}#mermaid-svg-XJ06xylSuyIkYSH4 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 .cluster-label text{fill:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 .cluster-label span{color:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 .cluster-label span p{background-color:transparent;}#mermaid-svg-XJ06xylSuyIkYSH4 .label text,#mermaid-svg-XJ06xylSuyIkYSH4 span{fill:#333;color:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 .node rect,#mermaid-svg-XJ06xylSuyIkYSH4 .node circle,#mermaid-svg-XJ06xylSuyIkYSH4 .node ellipse,#mermaid-svg-XJ06xylSuyIkYSH4 .node polygon,#mermaid-svg-XJ06xylSuyIkYSH4 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-XJ06xylSuyIkYSH4 .rough-node .label text,#mermaid-svg-XJ06xylSuyIkYSH4 .node .label text,#mermaid-svg-XJ06xylSuyIkYSH4 .image-shape .label,#mermaid-svg-XJ06xylSuyIkYSH4 .icon-shape .label{text-anchor:middle;}#mermaid-svg-XJ06xylSuyIkYSH4 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-XJ06xylSuyIkYSH4 .rough-node .label,#mermaid-svg-XJ06xylSuyIkYSH4 .node .label,#mermaid-svg-XJ06xylSuyIkYSH4 .image-shape .label,#mermaid-svg-XJ06xylSuyIkYSH4 .icon-shape .label{text-align:center;}#mermaid-svg-XJ06xylSuyIkYSH4 .node.clickable{cursor:pointer;}#mermaid-svg-XJ06xylSuyIkYSH4 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-XJ06xylSuyIkYSH4 .arrowheadPath{fill:#333333;}#mermaid-svg-XJ06xylSuyIkYSH4 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-XJ06xylSuyIkYSH4 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-XJ06xylSuyIkYSH4 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-XJ06xylSuyIkYSH4 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-XJ06xylSuyIkYSH4 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-XJ06xylSuyIkYSH4 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-XJ06xylSuyIkYSH4 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-XJ06xylSuyIkYSH4 .cluster text{fill:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 .cluster span{color:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-XJ06xylSuyIkYSH4 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-XJ06xylSuyIkYSH4 rect.text{fill:none;stroke-width:0;}#mermaid-svg-XJ06xylSuyIkYSH4 .icon-shape,#mermaid-svg-XJ06xylSuyIkYSH4 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-XJ06xylSuyIkYSH4 .icon-shape p,#mermaid-svg-XJ06xylSuyIkYSH4 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-XJ06xylSuyIkYSH4 .icon-shape .label rect,#mermaid-svg-XJ06xylSuyIkYSH4 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-XJ06xylSuyIkYSH4 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-XJ06xylSuyIkYSH4 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-XJ06xylSuyIkYSH4 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} DPO 路径
chosen/rejected 数据
policy vs ref
直接优化偏好 odds
经典 RLHF
Human Preference
Reward Model
PPO Actor+Critic


七、显存与硬件:RTX 3090 / 12GB+ 怎么配?

7.1 显存粗算(64M Dense,bf16 训练)

组成部分 量级
模型权重 64M × 2B ≈ 128MB
优化器 AdamW ~2× 参数量级(fp32 状态)
激活值 ∝ batch × seq_len × hidden × layers(主导项)
DPO ref 前向 额外一份 ref 权重(冻结,无优化器)

64M 很小,瓶颈几乎总在激活与 seq_len,不在参数量。

7.2 推荐配置表

显卡 Pretrain SFT DPO 备注
RTX 3090 24GB 默认 batch 32 可跑 默认较宽裕 batch 4 舒适 官方基准卡
RTX 3060 12GB batch 8~16,accumulation_steps 加大 batch 4~8 batch 1~2,seq≤512 能跑,需调参
RTX 4060 Ti 16GB 中等 batch 中等 batch 2~4 性价比较稳
笔记本 8GB 仅建议推理 eval_llm.py 训练易 OOM 不推荐 用云 3090

7.3 12GB 保命参数模板

bash 复制代码
# Pretrain @ 12GB
python train_pretrain.py \
  --batch_size 8 \
  --accumulation_steps 32 \
  --max_seq_len 512 \
  --dtype bfloat16 \
  --num_workers 4

# SFT @ 12GB
python train_full_sft.py \
  --batch_size 4 \
  --accumulation_steps 16 \
  --max_seq_len 512

# DPO @ 12GB(双前向,最吃显存)
python train_dpo.py \
  --batch_size 1 \
  --accumulation_steps 8 \
  --max_seq_len 384 \
  --dtype bfloat16

原则

  1. 先降 batch_size,用 accumulation_steps 补有效 batch
  2. 再降 max_seq_len(mini 数据 768→512 往往仍可收敛)
  3. 优先 bfloat16(3090/40 系支持;老卡用 float16 + GradScaler)
  4. DPO 阶段最吃显存:降低 seq 比降 lr 更优先

7.4 官方训练耗时参考(3090,1 epoch)

模型 Pretrain mini SFT mini DPO
minimind-3 64M ≈1.21h ≈1.10h 约 1h 量级
合计 Zero 路线 ≈2.3h / ≈3¥ 可选 +1h

八、常见报错与踩坑清单

8.1 CUDA OOM

现象torch.cuda.OutOfMemoryError

处理顺序

text 复制代码
1. 减半 batch_size
2. max_seq_len 768 → 512 → 384
3. accumulation_steps 翻倍(保持有效 batch)
4. num_workers 降到 2~4(有时 dataloader 也占内存)
5. DPO:batch_size=1 是常态
6. 关闭 --use_compile 1(编译期偶发额外显存)

8.2 Windows pyarrow / datasets DLL 冲突

源码注释(issue #771):

python 复制代码
import datasets  # noqa: F401 # Windows pyarrow/torch DLL conflict workaround

建议 :优先 WSL2 / Linux 训练;Windows 需对齐 Python、torch、pyarrow 版本。

8.3 续训找不到 checkpoint

bash 复制代码
python train_pretrain.py --from_resume 1

检查 ../checkpoints 下是否有对应 save_weight 的中间状态。

8.4 SFT 正常但对话很蠢

  • 是否用错权重:eval_llm.py --weight full_sft
  • Pretrain 是否充分(loss 还高就进 SFT)
  • max_seq_len 截断是否把答案截没

8.5 DPO 之后「变傻 / 只会拒绝」

原因 对策
learning_rate 太大 降到 ≤5e-8
beta 过大 试 0.1~0.15
epoch 过多 DPO 1 epoch 往往够
偏好数据质量差 rejected 太极端会学成「万事拒绝」
未基于 SFT --from_weight full_sft 必须

8.6 Loss 为 NaN

  • 检查数据是否有空文本、异常字符
  • grad_clip=1.0 已默认开启
  • 学习率减半
  • MoE 时关注 aux_loss 是否过大

8.7 中文乱码 / 模板错误

  • 确认使用仓库自带 tokenizer
  • SFT 数据需符合 conversations schema
  • Tool Call 样本不要经 pre_processing_chat 误删 tools 字段

九、2 小时最小复现路线(抄作业版)

bash 复制代码
# 0. 环境
conda create -n minimind python=3.10 -y && conda activate minimind
git clone https://github.com/jingyaogong/minimind.git && cd minimind
pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple

# 1. 数据(从 HuggingFace 或 README 链接下载到 dataset/)
#    需要:pretrain_t2t_mini.jsonl, sft_t2t_mini.jsonl, (可选)dpo.jsonl

# 2. Pretrain (~1.2h @ 3090)
cd trainer
python train_pretrain.py \
  --data_path ../dataset/pretrain_t2t_mini.jsonl \
  --max_seq_len 768 --epochs 1

# 3. SFT (~1.1h @ 3090)
python train_full_sft.py \
  --from_weight pretrain \
  --data_path ../dataset/sft_t2t_mini.jsonl \
  --max_seq_len 768 --epochs 1

# 4. 对话测试
python eval_llm.py --weight full_sft

# 5. (可选)DPO 对齐
python train_dpo.py --from_weight full_sft --data_path ../dataset/dpo.jsonl
python eval_llm.py --weight dpo

验收

  • Pretrain:eval_llm.py --weight pretrain 能续写常识
  • SFT:能按 user/assistant 模板多轮聊天
  • DPO:同一问题上更礼貌、更符合 chosen 风格(主观对比)

十、读完源码,你应带走的 5 个「真懂」检查点

  1. Pretrain 全序列 loss;SFT 只 mask 学 assistant ------ 看 lm_dataset.pygenerate_labels
  2. DPO = policy 与 ref 的 logprob 差分 + sigmoid ------ 看 train_dpo.pydpo_loss
  3. 64M = 768×8 层 + 6400 词表 + GQA,不是魔法
  4. 显存瓶颈在 seq_len × batch,不在 64M 参数
  5. 对齐分两层:SFT 学能力/格式,DPO 学偏好;别指望 DPO 单独「变聪明」

十一、延伸:MiniMind 还藏了什么?

读完主线后,可按兴趣继续:

脚本 内容
train_lora.py 从 0 实现 LoRA,无 peft
train_ppo.py / train_grpo.py RLAIF,可验证奖励
train_distillation.py 白盒蒸馏
train_agent.py Tool Use / Agentic RL
scripts/web_demo.py Streamlit 对话
scripts/convert_model.py 转 GGUF 等

仓库价值不仅是「炼出 64M」,而是 把 LLM 训练全链路拆成可读、可改、可跑的 PyTorch 脚本------这才是适合中国开发者「搞懂再上大模型」的教材级项目。


浓缩总结(公众号结尾卡)

text 复制代码
MiniMind = 64M Dense 小模型 + 原生 PyTorch 全链路

四步炼成:
  数据 jsonl → BPE 分词 → Pretrain 接龙 → SFT 只学回答 → DPO 偏好对齐

硬件:
  3090 约 2h/3¥ 跑通 Zero;12GB 降 batch/seq 也能炼

对齐:
  DPO 看 train_dpo.py 三十行;lr≤5e-8,防遗忘

真懂标准:
  能解释 labels=-100 在哪一步、DPO 的 ref 为何冻结

参考资源


相关推荐
IT_陈寒1 小时前
Redis持久化配置漏了这一步,线上数据丢了5小时
前端·人工智能·后端
大力财经1 小时前
抖音生活服务品牌零售行业峰会在杭州举办,探索线下生意新增量
大数据·人工智能·区块链
xsd202411181 小时前
检测视觉大模型全景解析:从Grounding DINO到Molmo,AI如何“指哪打哪“
人工智能
咖啡星人k1 小时前
2026 智能体安全进阶:把注入和越权写进SPEC,MonkeyCode 云端跑通
人工智能·安全·机器学习
sel_91 小时前
深度学习激活函数详解:从 Sigmoid、Tanh、ReLU 到 GELU、SiLU、Mish,一文掌握所有常用激活函数
人工智能·深度学习
智擎GEO1 小时前
中山GEO优化哪个靠谱
人工智能·python
liliangcsdn1 小时前
多空价差收益序列汇总统计的分析和代码示例
人工智能·算法·机器学习
aneasystone本尊1 小时前
学习大模型推理的 KV Cache
人工智能
单词记忆方法研究1 小时前
专项突破词汇实操教程与要点解析
人工智能