chenbhao commited on
Commit
d58698c
·
1 Parent(s): 64d62ee

reorganize configs into subdirs lm/vlm/vam, remove runs/, add interview docs for pretrain & full_sft

Browse files
README.md CHANGED
@@ -67,18 +67,15 @@ src/
67
  │ └── checkpoint.py # checkpoint 读写辅助
68
  ├── serve/ # 实时语音会话(SileroVAD / RealtimeSession)
69
  configs/
70
- ├── lm.yaml # 纯文本训练配置
71
- ├── lm_moe.yaml # 纯文本 MoE 训练配置
 
 
72
  ├── vlm.yaml # 视觉多模态训练配置
73
- ├── vlm_moe.yaml # 视觉多模态 MoE 训练配置
74
- ├── vam.yaml # 模态训练配置
75
- ├── vam_moe.yaml # 全模态 MoE 训练配置
76
  └── tokenizer/ # tokenizer.json / tokenizer_config.json
77
- runs/ # 根目录可直接运行的训练启动脚本(默认加载对应 configs/*.yaml)
78
- ├── train_lm.sh # bash runs/train_lm.sh -> configs/lm.yaml
79
- ├── train_vlm.sh # bash runs/train_vlm.sh -> configs/vlm.yaml
80
- ├── train_vam.sh # bash runs/train_vam.sh -> configs/vam.yaml
81
- └── train_tokenizer.sh # bash runs/train_tokenizer.sh -> 训练 tokenizer,保存到 checkpoint/tokenizer/
82
  scripts/ # 推理 / 服务 / 转换脚本
83
  ├── eval_llm.py # 命令行推理与对话
84
  ├── eval_vlm.py # 视觉多模态推理
@@ -102,20 +99,79 @@ pip install -e ".[rl,serve,demo]"
102
 
103
  ### 训练(YAML 驱动)
104
 
105
- 根目录 `runs/` 提供可直接运行的 `.sh` 启动脚本,默认加载 `configs/` 下对应的 YAML,
106
- 也可通过 `--config` 指定其它配置,任意 CLI 参数都能覆盖 YAML 中的默认值:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
107
 
108
  ```bash
109
- bash runs/train_tokenizer.sh --data_path dataset/sft_t2t_mini.jsonl \
110
- --vocab_size 6400 \
111
- --checkpoint_dir ./checkpoint \
112
- --no_eval
113
 
114
- bash runs/train_lm.sh # 使用 configs/lm.yaml
115
- bash runs/train_vlm.sh --config configs/vlm_moe.yaml
116
- bash runs/train_vam.sh --epochs 5 # 在 vam.yaml 基础上覆盖单字段
117
 
 
 
 
118
 
 
 
 
119
  ```
120
 
121
  ### 推理 / 对话
 
67
  │ └── checkpoint.py # checkpoint 读写辅助
68
  ├── serve/ # 实时语音会话(SileroVAD / RealtimeSession)
69
  configs/
70
+ ├── lm_full_sft.yaml # 纯文本 SFT 训练配置
71
+ ├── lm_full_sft_moe.yaml # 纯文本 MoE SFT 训练配置
72
+ ├── lm_pretrain.yaml # 纯文本预训练配置
73
+ ├── lm_pretrain_moe.yaml # 纯文本 MoE 预训练配置
74
  ├── vlm.yaml # 视觉多模态训练配置
75
+ ├── lm/ # 纯文本 LM 配置
76
+ ├── vlm/ # 视觉多模态 VLM 配置
77
+ ├── vam/ # 全模态 VAM 配置
78
  └── tokenizer/ # tokenizer.json / tokenizer_config.json
 
 
 
 
 
79
  scripts/ # 推理 / 服务 / 转换脚本
80
  ├── eval_llm.py # 命令行推理与对话
81
  ├── eval_vlm.py # 视觉多模态推理
 
99
 
100
  ### 训练(YAML 驱动)
101
 
102
+ 训练入口统一为 `python -m trainers.<包>.<脚本>`通过 `--config` 指定 YAML 配置
103
+ 任意 CLI 参数都能覆盖 YAML 中的默认值:
104
+
105
+ #### 纯文本 LM
106
+
107
+ ```bash
108
+ # 预训练(从头训练语言模型)
109
+ python -m trainers.lm.pretrain --config configs/lm/lm_pretrain.yaml
110
+
111
+ # 全量 SFT(以预训练权重初始化,指令微调)
112
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft.yaml
113
+
114
+ # 训练 tokenizer(学习用,MiniMind 已自带)
115
+ python -m trainers.lm.train_tokenizer --data_path dataset/sft_t2t_mini.jsonl \
116
+ --vocab_size 6400 \
117
+ --checkpoint_dir ./checkpoint \
118
+ --no_eval
119
+
120
+ # LoRA 微调
121
+ python -m trainers.lm.lora_sft --config configs/lm/lm_full_sft.yaml
122
+
123
+ # DPO / PPO / GRPO 偏好对齐
124
+ python -m trainers.lm.dpo --config configs/lm/lm_full_sft.yaml
125
+ python -m trainers.lm.ppo --config configs/lm/lm_full_sft.yaml
126
+ python -m trainers.lm.grpo --config configs/lm/lm_full_sft.yaml
127
+
128
+ # 知识蒸馏
129
+ python -m trainers.lm.distill --teacher <teacher_path> --config configs/lm/lm_full_sft.yaml
130
+
131
+ # MoE 变体
132
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft_moe.yaml
133
+ python -m trainers.lm.pretrain --config configs/lm/lm_pretrain_moe.yaml
134
+
135
+ # Mini 变体(快速验证用,h=128, L=4, ~14min pretrain)
136
+ python -m trainers.lm.pretrain --config configs/lm/lm_pretrain_mini.yaml
137
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft_mini.yaml
138
+ ```
139
+
140
+ #### 视觉多模态 VLM
141
+
142
+ ```bash
143
+ # 预训练(视觉模态对齐)
144
+ python -m trainers.vlm.pretrain --config configs/vlm/vlm.yaml
145
+
146
+ # 全量 SFT
147
+ python -m trainers.vlm.full_sft --config configs/vlm/vlm.yaml
148
+ python -m trainers.vlm.full_sft --config configs/vlm/vlm_moe.yaml # MoE 变体
149
+ ```
150
+
151
+ #### 全模态 VAM(文本 + 视觉 + 语音)
152
+
153
+ ```bash
154
+ # 全量 SFT
155
+ python -m trainers.vam.full_sft --config configs/vam/vam.yaml
156
+ python -m trainers.vam.full_sft --config configs/vam/vam_moe.yaml # MoE 变体
157
+ ```
158
+
159
+ #### 常用覆盖参数
160
 
161
  ```bash
162
+ # 覆盖任意 YAML 字段
163
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft.yaml \
164
+ --epochs 3 --batch_size 8 --learning_rate 5e-6
 
165
 
166
+ python -m trainers.vam.full_sft --config configs/vam/vam.yaml --epochs 5
 
 
167
 
168
+ # 指定 checkpoint 目录
169
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft.yaml \
170
+ --save_dir checkpoint/my_exp
171
 
172
+ # 多卡 DDP 训练(torchrun)
173
+ torchrun --nproc_per_node=4 -m trainers.lm.pretrain --config configs/lm/lm_pretrain.yaml
174
+ torchrun --nproc_per_node=4 -m trainers.lm.full_sft --config configs/lm/lm_full_sft.yaml
175
  ```
176
 
177
  ### 推理 / 对话
configs/{lm.yaml → lm/lm_full_sft.yaml} RENAMED
@@ -19,8 +19,9 @@ train:
19
  save_interval: 1000
20
  log_interval: 100
21
  from_weight: pretrain # none / pretrain / full_sft
 
22
  from_resume: 0
23
 
24
  paths:
25
  save_dir: checkpoint/lm
26
- data_path: dataset/sft.jsonl
 
19
  save_interval: 1000
20
  log_interval: 100
21
  from_weight: pretrain # none / pretrain / full_sft
22
+ model_dir: checkpoint/lm_pretrain
23
  from_resume: 0
24
 
25
  paths:
26
  save_dir: checkpoint/lm
27
+ data_path: dataset/sft_t2t_mini.jsonl
configs/lm/lm_full_sft_mini.yaml ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # LM SFT 配置(小模型快速微调,使用 lm_pretrain_mini 的预训练权重)
2
+ model:
3
+ hidden_size: 128
4
+ num_hidden_layers: 4
5
+ use_moe: 0
6
+ vocab_size: 6400
7
+ max_seq_len: 768
8
+ num_attention_heads: 8
9
+ num_key_value_heads: 4
10
+ dropout: 0.0
11
+
12
+ train:
13
+ epochs: 1
14
+ batch_size: 32
15
+ learning_rate: 1.0e-5
16
+ accumulation_steps: 1
17
+ grad_clip: 1.0
18
+ dtype: bfloat16
19
+ save_interval: 500
20
+ log_interval: 50
21
+ from_weight: pretrain
22
+ model_dir: checkpoint/lm_pretrain_mini
23
+ from_resume: 0
24
+
25
+ paths:
26
+ save_dir: checkpoint/lm_full_sft_mini
27
+ data_path: dataset/sft_t2t_mini.jsonl
configs/{lm_moe.yaml → lm/lm_full_sft_moe.yaml} RENAMED
File without changes
configs/lm/lm_pretrain.yaml ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # LM 预训练配置(纯文本语言模型,从头训练)
2
+ model:
3
+ hidden_size: 512
4
+ num_hidden_layers: 8
5
+ use_moe: 0 # 1 启用 MoE
6
+ vocab_size: 6400
7
+ max_seq_len: 340
8
+ num_attention_heads: 8
9
+ num_key_value_heads: 4
10
+ dropout: 0.0
11
+
12
+ train:
13
+ epochs: 2
14
+ batch_size: 32
15
+ learning_rate: 5.0e-4
16
+ accumulation_steps: 8
17
+ grad_clip: 1.0
18
+ dtype: bfloat16
19
+ save_interval: 1000
20
+ log_interval: 100
21
+ from_weight: none # 从头开始预训练
22
+ from_resume: 0
23
+
24
+ paths:
25
+ save_dir: checkpoint/lm_pretrain
26
+ data_path: dataset/pretrain_t2t_mini.jsonl
configs/lm/lm_pretrain_mini.yaml ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # LM 预训练配置(小模型快速验证,~30分钟跑完)
2
+ model:
3
+ hidden_size: 128
4
+ num_hidden_layers: 4
5
+ use_moe: 0
6
+ vocab_size: 6400
7
+ max_seq_len: 256
8
+ num_attention_heads: 8
9
+ num_key_value_heads: 4
10
+ dropout: 0.0
11
+
12
+ train:
13
+ epochs: 1
14
+ batch_size: 96
15
+ learning_rate: 5.0e-4
16
+ accumulation_steps: 1
17
+ grad_clip: 1.0
18
+ dtype: bfloat16
19
+ save_interval: 500
20
+ log_interval: 50
21
+ from_weight: none
22
+ from_resume: 0
23
+
24
+ paths:
25
+ save_dir: checkpoint/lm_pretrain_mini
26
+ data_path: dataset/pretrain_t2t_mini.jsonl
configs/lm/lm_pretrain_moe.yaml ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # LM-MoE 预训练配置(纯文本语言模型,启用 MoE,从头训练)
2
+ model:
3
+ hidden_size: 512
4
+ num_hidden_layers: 8
5
+ use_moe: 1 # 1 启用 MoE
6
+ num_experts: 4
7
+ num_experts_per_tok: 1
8
+ moe_intermediate_size: 1408
9
+ norm_topk_prob: true
10
+ router_aux_loss_coef: 5.0e-4
11
+ vocab_size: 6400
12
+ max_seq_len: 340
13
+ num_attention_heads: 8
14
+ num_key_value_heads: 4
15
+ dropout: 0.0
16
+
17
+ train:
18
+ epochs: 2
19
+ batch_size: 32
20
+ learning_rate: 5.0e-4
21
+ accumulation_steps: 8
22
+ grad_clip: 1.0
23
+ dtype: bfloat16
24
+ save_interval: 1000
25
+ log_interval: 100
26
+ from_weight: none # 从头开始预训练
27
+ from_resume: 0
28
+
29
+ paths:
30
+ save_dir: checkpoint/lm_pretrain_moe
31
+ data_path: dataset/pretrain_t2t_mini.jsonl
configs/{vam.yaml → vam/vam.yaml} RENAMED
File without changes
configs/{vam_moe.yaml → vam/vam_moe.yaml} RENAMED
File without changes
configs/{vlm.yaml → vlm/vlm.yaml} RENAMED
File without changes
configs/{vlm_moe.yaml → vlm/vlm_moe.yaml} RENAMED
File without changes
docs/architecture/overview.md CHANGED
@@ -41,10 +41,9 @@ input_ids
41
  ## 4. 训练入口
42
 
43
  ```
44
- runs/
45
- ├── lm.sh → python -m trainers.lm.full_sft --config configs/lm.yaml
46
- ├── vlm.sh → python -m trainers.vlm.full_sft --config configs/vlm.yaml
47
- └── vam.sh → python -m trainers.vam.full_sft --config configs/vam.yaml
48
  ```
49
 
50
  配置读取:`utils.training.apply_config(parser, default_config)` 把 YAML 的
 
41
  ## 4. 训练入口
42
 
43
  ```
44
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft.yaml
45
+ python -m trainers.vlm.full_sft --config configs/vlm/vlm.yaml
46
+ python -m trainers.vam.full_sft --config configs/vam/vam.yaml
 
47
  ```
48
 
49
  配置读取:`utils.training.apply_config(parser, default_config)` 把 YAML 的
docs/interview/full_sft.md ADDED
@@ -0,0 +1,528 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 面试:全量微调(Full SFT)深度
2
+
3
+ > 本仓库 `src/trainers/lm/full_sft.py`,对应 config:`configs/lm/lm_full_sft.yaml`
4
+
5
+ ## 0. 整体流程
6
+
7
+ ```
8
+ 数据集 (sft_t2t_mini.jsonl)
9
+ │ 60,000 条多轮对话
10
+
11
+ SFTDataset
12
+ ├── 解析 JSON 对话格式
13
+ ├── 拼接多轮对话(含 <s>system</s> <s>user</s> <s>assistant</s> 标记)
14
+ ├── input_ids:全部 token 的 id
15
+ ├── labels:prompt 段 → -100,assistant 段 → input_ids
16
+ └── 截断 / 填充到 max_seq_len
17
+
18
+
19
+ DataLoader → model(input_ids, labels)
20
+
21
+ ├── 🔥 加载预训练权重(from_weight=pretrain)
22
+ │ └── model_dir 指定权重源目录
23
+
24
+ ├── 全量参数更新(所有参数参与训练)
25
+
26
+ └── Loss = CE(shift_logits, shift_labels)
27
+ 只计算 labels != -100 的位置
28
+
29
+
30
+ AdamW + Cosine LR → 保存 checkpoint
31
+ ```
32
+
33
+ ---
34
+
35
+ ## Q1. SFTDataset 如何处理多轮对话?
36
+
37
+ ### 输入格式
38
+
39
+ ```json
40
+ {
41
+ "conversations": [
42
+ {"from": "system", "value": "你是一个 AI 助手"},
43
+ {"from": "user", "value": "你好"},
44
+ {"from": "assistant", "value": "你好!有什么可以帮你的?"},
45
+ {"from": "user", "value": "什么是 LLM?"},
46
+ {"from": "assistant", "value": "LLM 是大型语言模型..."}
47
+ ]
48
+ }
49
+ ```
50
+
51
+ ### 编码拼接(`src/dataset/sft.py`)
52
+
53
+ ```python
54
+ def _build_conversation(self, conversation):
55
+ input_ids, labels = [], []
56
+ for turn in conversation:
57
+ speaker = turn['from']
58
+ text = turn['value']
59
+ # 用特殊标记包裹每条发言
60
+ tokens = self.tokenizer.encode(f'<s>{speaker}</s>\n{text}\n')
61
+ if speaker == 'assistant':
62
+ input_ids += tokens
63
+ labels += tokens # 回复段参与损失
64
+ else:
65
+ input_ids += tokens
66
+ labels += [-100] * len(tokens) # prompt 段忽略损失
67
+ return input_ids[:self.max_length], labels[:self.max_length]
68
+ ```
69
+
70
+ ### 为什么 prompt 段 label 置 -100?
71
+
72
+ - `nn.CrossEntropyLoss(ignore_index=-100)` 自动忽略这些位置的损失
73
+ - **只学习回复内容**:模型只需要学会生成 assistant 的回答,不需要拟合用户的输入
74
+ - **保持 prompt 长度灵活性**:不需要 mask attention,只是不算损失
75
+
76
+ ### 极端情况
77
+
78
+ 如果整条对话 `max_seq_len` 截断后只包含 prompt 段(user 说的话),那么所有的 labels 都是 -100,loss = 0。这种情况虽然罕见但需要注意——数据预处理时应该过滤掉这种样本。
79
+
80
+ > 面试点:如果一条对话全部被截断为 prompt 端怎么办?→ loss=0,该样本对训练无贡献;需要在数据预处理时过滤或截断时尽量保留 assistant 段
81
+
82
+ ---
83
+
84
+ ## Q2. 加载预训练权重机制
85
+
86
+ ### 权重加载流程
87
+
88
+ ```
89
+ 1. from_weight=pretrain
90
+
91
+ 2. 确定权重路径:
92
+ model_dir / from_weight_hidden_size.pth
93
+ → checkpoint/lm_pretrain/pretrain_512.pth
94
+
95
+ 3. torch.load(..., map_location=device)
96
+
97
+ 4. model.load_state_dict(weights, strict=False)
98
+
99
+ 5. 选择不加载 lm_head.weight(可选)
100
+ ```
101
+
102
+ ### 代码实现
103
+
104
+ ```python
105
+ # src/utils/training.py:145-168
106
+ def init_model(config, save_dir, weight_type='pretrain', model_dir=None):
107
+ if weight_type != 'none':
108
+ weight_dir = model_dir or save_dir # model_dir 优先
109
+ weight_path = f'{weight_dir}/{weight_type}_{config.hidden_size}.pth'
110
+ weights = torch.load(weight_path, map_location=device)
111
+ # 可选择跳过 lm_head(如新增词表时)
112
+ # weights = {k: v for k, v in weights.items() if 'lm_head' not in k}
113
+ model.load_state_dict(weights, strict=False)
114
+ ```
115
+
116
+ ### model_dir vs save_dir
117
+
118
+ | 参数 | 作用 | 默认值 |
119
+ |---|---|---|
120
+ | `save_dir` | checkpoint 写入目录 | 配置中的 `paths.save_dir` |
121
+ | `model_dir` | checkpoint 读取目录 | 未设置时 = save_dir |
122
+
123
+ 为什么要分开?
124
+ - pretrain 权重存在 `checkpoint/lm_pretrain/`
125
+ - SFT 权重存在 `checkpoint/lm/`
126
+ - SFT 需要从 pretrain 目录**加载**,但**保存**到自己目录
127
+ - 没有 `model_dir` 时会在 `checkpoint/lm/` 下找 `pretrain_512.pth`,找不到
128
+
129
+ ### 不同 from_weight 的语义
130
+
131
+ | from_weight | 用途 | 加载的文件 |
132
+ |---|---|---|
133
+ | `none` | 从头训练 | 不加载(随机初始化) |
134
+ | `pretrain` | pretrain → SFT 持续训练 | `pretrain_{h}.pth` |
135
+ | `full_sft` | SFT 继续训练/增量 SFT | `full_sft_{h}.pth` |
136
+
137
+ ### strict=False 的注意事项
138
+
139
+ - 允许权重文件和模型结构**不完全一致**
140
+ - 常见的 mismatch 来源:
141
+ 1. embedding 和 lm_head 的 weight tying(两个 key 映射到同一个参数字典)
142
+ 2. 词表大小变化(新增 special token 后加载旧权重)
143
+ 3. 模型架构微调(如增减层数)
144
+ - `strict=False` 会静默忽略多出的键(缺少的键会随机初始化)
145
+
146
+ > 面试点:strict=False 实际可能导致模型部分随机初始化而不报错,如何确保没有遗漏?→ 加载后比较 `model.state_dict().keys()` 和 `weights.keys()`,打印 missing_keys 和 unexpected_keys
147
+
148
+ ---
149
+
150
+ ## Q3. SFT 训练流程详解
151
+
152
+ ### 训练循环(`src/trainers/lm/full_sft.py`)
153
+
154
+ ```python
155
+ def train_epoch(epoch, model, loader, optimizer, scheduler, scaler, args):
156
+ model.train()
157
+ total_loss = 0
158
+ for step, (input_ids, labels) in enumerate(loader):
159
+ input_ids = input_ids.cuda()
160
+ labels = labels.cuda()
161
+
162
+ with autocast_ctx:
163
+ logits = model(input_ids)
164
+ loss = cross_entropy(logits[..., :-1, :].contiguous(),
165
+ labels[..., 1:].contiguous(),
166
+ ignore_index=-100)
167
+ loss = loss / args.accumulation_steps
168
+
169
+ scaler.scale(loss).backward()
170
+
171
+ if (step + 1) % args.accumulation_steps == 0:
172
+ scaler.unscale_(optimizer)
173
+ clip_grad_norm_(model.parameters(), args.grad_clip)
174
+ scaler.step(optimizer)
175
+ scaler.update()
176
+ optimizer.zero_grad(set_to_none=True)
177
+ scheduler.step()
178
+
179
+ total_loss += loss.item()
180
+ ```
181
+
182
+ ### 与 Pretrain 训练循环的区别
183
+
184
+ | | Pretrain | Full SFT |
185
+ |---|---|---|
186
+ | 是否加载预训练权重 | 否(from_weight=none) | 是(from_weight=pretrain) |
187
+ | 学习率 | 5e-4(从零学习) | 1e-5(微调,小 lr) |
188
+ | 训练轮数 | 2 | 2 |
189
+ | 数据集 | 纯文本 1.27M 条 | 指令对话 60K 条 |
190
+ | loss 计算 | 全 token | 仅 assistant 段 |
191
+ | weight_decay | 有 | 有(默认 0.1) |
192
+
193
+ ### 为什么 SFT 学习率要小?
194
+
195
+ - 预训练权重已经学到很好的语言表示
196
+ - 大学习率会破坏/覆盖预训练知识(灾难性遗忘)
197
+ - 目标是在已有知识基础上"微调"指令跟随能力
198
+ - 一般 pretrain LR : SFT LR ≈ 10:1 ~ 50:1
199
+
200
+ ---
201
+
202
+ ## Q4. 灾难性遗忘(Catastrophic Forgetting)
203
+
204
+ ### 什么灾难性遗忘?
205
+
206
+ 模型在 SFT 阶段学会生成对话回复的同时,**丢失**了在预训练阶段学到的通用知识(如常识推理、知识问答能力)。
207
+
208
+ ### 为什么 SFT 会导致遗忘?
209
+
210
+ ```
211
+ 预训练: 语言 P(x₁...xₙ) ← 通用分布
212
+ SFT: 条件 P(回复|指令) ← 狭窄分布
213
+ ──────────────────────→
214
+ 训练分布偏移 → 覆盖预训练权重
215
+ ```
216
+
217
+ ### 缓解策略
218
+
219
+ 1. **小学习率**:1e-5 比 5e-4 小 50 倍,梯度更新量小
220
+ 2. **少轮数**:1-2 轮足够,更多轮次会导致过拟合和遗忘
221
+ 3. **保留预训练数据**:混合 SFT 数据 + 10-20% 预训练数据(本仓库未实现)
222
+ 4. **EWC / LwF**:正则化方法,限制重要参数的大幅更新(本仓库未实现)
223
+ 5. **LoRA**:增量微调,冻结原权重(本仓库另有 `src/trainers/lm/lora_sft.py`)
224
+
225
+ > 面试点:什么情况下灾难性遗忘最严重?→ 大量 SFT 数据 + 高学习率 + 多轮训练 + 领域单一的数据集
226
+
227
+ ---
228
+
229
+ ## Q5. SFT 训练问题诊断
230
+
231
+ ### Loss 正常值范围
232
+
233
+ | 阶段 | Loss | 说明 |
234
+ |---|---|---|
235
+ | 初始(第 1 步) | ~7.0-8.5 | 刚加载 pretrain 权重,但换数据集后分布不同 |
236
+ | 收敛 | ~1.5-2.5 | 模型学会生成合理回复 |
237
+ | 过拟合 | < 1.0 | 训练 loss 极低但生成质量差(记忆而不是泛化) |
238
+
239
+ ### Loss 异常分析
240
+
241
+ ```
242
+ Loss 行为 可能原因 建议
243
+ ────────────────────────────────────────────────────────────────────
244
+ 初始 loss 极低 (<2.0) pretrain 数据集和 SFT 高度重叠 check 数据分布
245
+ loss 不下降 (<30 步) 学习率太小 / 模型冻结 check 梯度
246
+ loss 突增 学习率太大 / 梯度爆炸 减少 lr / 加强 grad_clip
247
+ loss 震荡 batch 太小 / lr 太高 调大 batch / 减小 lr
248
+ loss 降到 0.0 所有 label 为 -100 检查数据截断
249
+ val loss 上升但 train train loss 下降过拟合 early stopping / 正则化
250
+ ```
251
+
252
+ ---
253
+
254
+ ## Q6. 推理时生成差异
255
+
256
+ ### 训练 vs 推理行为对比
257
+
258
+ ```python
259
+ # 训练时
260
+ model.train()
261
+ logits = model(input_ids) # 全部序列
262
+ loss = cross_entropy(shift_logits, shift_labels) # 不采样,计算损失
263
+
264
+ # 推理时
265
+ model.eval()
266
+ generated = model.generate(input_ids, max_new_tokens=256, do_sample=True, temperature=0.7)
267
+ ```
268
+
269
+ ### 推理超参
270
+
271
+ | 参数 | 作用 | SFT 推荐值 |
272
+ |---|---|---|
273
+ | `do_sample` | 是否采样(否则 greedy) | `True` |
274
+ | `temperature` | 采样温度,越高越随机 | 0.7-0.9 |
275
+ | `top_k` | 只从前 K 个 token 采样 | 40-50 |
276
+ | `top_p` | 核采样(累积概率 p) | 0.9 |
277
+ | `repetition_penalty` | 重复惩罚 | 1.05-1.15 |
278
+
279
+ ### 温度对比
280
+
281
+ ```
282
+ Temperature=0.1: "今天天气真好,我们去散步吧。"
283
+ Temperature=0.7: "今天天气真好,要不出去走走?"
284
+ Temperature=1.5: "天气不错,散步散步散步吧...哦不对不对呵呵呵"
285
+ ```
286
+
287
+ - 温度太低 → 输出机械、重复
288
+ - 温度太高 → 输出发散、语无伦次
289
+ - 0.7 是创造性 + 连贯性的良好平衡点
290
+
291
+ ---
292
+
293
+ ## Q7. YAML 配置详解
294
+
295
+ ### 结构说明
296
+
297
+ ```yaml
298
+ model:
299
+ hidden_size: 512 # 模型容量:影响参数量和激活值大小
300
+ num_hidden_layers: 8 # Transformer 层数
301
+ use_moe: 0 # MoE 开关
302
+ vocab_size: 6400 # 词表大小
303
+ max_seq_len: 768 # SFT 通常需要更长上下文
304
+ num_attention_heads: 8 # Q 头数
305
+ num_key_value_heads: 4 # K/V 头数(GQA,4:8=2x 压缩)
306
+ dropout: 0.0 # SFT 一般不加 dropout
307
+
308
+ train:
309
+ epochs: 2
310
+ batch_size: 16 # 受显存限制
311
+ learning_rate: 1.0e-5 # 微调用小 lr
312
+ accumulation_steps: 1
313
+ grad_clip: 1.0
314
+ dtype: bfloat16
315
+ save_interval: 1000
316
+ log_interval: 100
317
+ from_weight: pretrain # 加载预训练权重
318
+ model_dir: checkpoint/lm_pretrain # 预训练权重来源目录
319
+ from_resume: 0
320
+
321
+ paths:
322
+ save_dir: checkpoint/lm # 训练产物存放目录
323
+ data_path: dataset/sft.jsonl
324
+ ```
325
+
326
+ ### max_seq_len 为什么比 pretrain 大?
327
+
328
+ | | Pretrain | SFT |
329
+ |---|---|---|
330
+ | max_seq_len | 340 | 768 |
331
+ | 原因 | 预训练数据多为短文本(如 BERT 风格片段) | 多轮对话需要更多空间 |
332
+
333
+ 对话拼接后 token 数≈ sum of turns,通常比单篇文本长。
334
+
335
+ ---
336
+
337
+ ## Q8. AdamW 优化器
338
+
339
+ ### 与 Adam 的区别
340
+
341
+ ```python
342
+ # Adam: w_{t+1} = w_t - lr * m_hat / (sqrt(v_hat) + eps) # 无 weight decay
343
+ # AdamW: w_{t+1} = w_t - lr * (m_hat / (sqrt(v_hat) + eps) + λ*w_t)
344
+ # └─────────────────────────────┬──────────────┘
345
+ # └ weight decay 与梯度更新解耦
346
+ ```
347
+
348
+ Adam 将 weight decay 和 L2 正则化混在一起(L2 = 在 loss 上加 λ/2 × ||w||²),而 AdamW 将 weight decay 从自适应学习率中解耦出来。
349
+
350
+ ### 为什么 AdamW 更好?
351
+
352
+ | | Adam (L2) | AdamW |
353
+ |---|---|---|
354
+ | Decay 位置 | loss 函数中(对梯度贡献) | optimizer 更新时独立加 |
355
+ | 自适应影响 | decay 也被 m_hat/v_hat 缩放 | decay 不受影响 |
356
+ | 实际效果 | 大学习率下 decay 被自适应削弱 | 稳定的 decay 效果 |
357
+ | 业界标准 | 旧方法 | GPT/LLaMA 等现代模型标配 |
358
+
359
+ ### 本仓库配置
360
+
361
+ ```python
362
+ optimizer = AdamW(model.parameters(), lr=args.learning_rate, weight_decay=0.1)
363
+ ```
364
+
365
+ - `weight_decay=0.1` 是常见推荐值
366
+ - 一般不对 bias 和 norm 参数做 weight decay(但此项目未做区分)
367
+ - PyTorch 的 AdamW 默认 `betas=(0.9, 0.999)`, `eps=1e-8`
368
+
369
+ ---
370
+
371
+ ## Q9. SFT 评估方法
372
+
373
+ ### 评估维度
374
+
375
+ | 维度 | 评估方式 | 指标 |
376
+ |---|---|---|
377
+ | 指令跟随 | 人工/模型评估 | 是否按指令完成 |
378
+ | 生成质量 | 人工评分 | 连贯性/有用性/安全性 |
379
+ | 多样性 | 统计 | distinct-1/2, ngram 重复率 |
380
+ | 知识正确性 | 基准测试 | MMLU, CEval, CMMLU |
381
+
382
+ ### 本仓库评估脚本(`scripts/eval_llm.py`)
383
+
384
+ ```python
385
+ # 生成回复 + 即时交互评估
386
+ python scripts/eval_llm.py --config configs/lm/lm_full_sft.yaml --checkpoint /path/to/full_sft_512.pth
387
+ ```
388
+
389
+ ```
390
+ 生成评估结果(示例):
391
+ ────────────────────────────────────
392
+ User: 讲个笑话
393
+ Assistant: 为什么程序员总把万圣节和圣诞节搞混?
394
+ 因为 Oct 31 == Dec 25!
395
+ ────────────────────────────────────
396
+ User: 用 Python 写一个快速排序
397
+ Assistant: def quicksort(arr):
398
+ if len(arr) <= 1: return arr
399
+ pivot = arr[len(arr)//2]
400
+ left = [x for x in arr if x < pivot]
401
+ mid = [x for x in arr if x == pivot]
402
+ right = [x for x in arr if x > pivot]
403
+ return quicksort(left) + mid + quicksort(right)
404
+ ────────────────────────────────────
405
+ ```
406
+
407
+ ### 常见问题
408
+
409
+ - **回复过短**:「是的」「好的」→ 数据集过于简单或数据量不够
410
+ - **回复重复**:不断生成相同短语 → 温度太低或 repetition_penalty 太小
411
+ - **偏离主题**:模型开始乱说 → 训练不足或 temperature 太高
412
+ - **不能按格式输出**:要求 JSON 但输出自然语言 → 数据集中缺乏格式化示例
413
+
414
+ ---
415
+
416
+ ## Q10. SFT vs RLHF 的关系
417
+
418
+ ### SFT 的局限性
419
+
420
+ 1. **模仿而非优化**:SFT 只是让模型模仿人工回复分布,不是优化最终效果
421
+ 2. **暴露偏差**:训练时使用 teacher forcing(每步输入真实 token),推理时输入是自生成的 token,分布偏移
422
+ 3. **缺乏偏好对齐**:所有训练样本被视为同等正确,区分不出"好��答"和"更好回答"
423
+
424
+ ### RLHF 如何解决?
425
+
426
+ ```
427
+ SFT 阶段:模仿示范数据
428
+
429
+
430
+ Reward 模型训练:学习偏好排序
431
+
432
+
433
+ PPO 阶段:以 reward 为信号优化策略
434
+
435
+
436
+ 结果:模型知道什么"更好",不仅仅是"像什么"
437
+ ```
438
+
439
+ ### 本仓库的 RL 系列
440
+
441
+ - `src/trainers/lm/dpo.py`:Direct Preference Optimization(PPO 的简化替代)
442
+ - `src/trainers/lm/ppo.py`:Proximal Policy Optimization(标准 RLHF)
443
+ - `src/trainers/lm/grpo.py`:Group Relative Policy Optimization(DeepSeek 方案)
444
+ - `src/trainers/lm/distill.py`:知识蒸馏
445
+
446
+ > 面试点:SFT 和 RLHF 的核心区别是什么?→ SFT 是监督学习(模仿示范),RLHF 是从偏好信号中学习优化(区分好与更好)
447
+
448
+ ---
449
+
450
+ ## Q11. Teacher Forcing 与 Exposure Bias
451
+
452
+ ### Teacher Forcing
453
+
454
+ ```python
455
+ # 训练时:每次输入真实 token
456
+ for t in range(seq_len):
457
+ logit = model(input_ids[:, :t+1])
458
+ loss = CE(logit[:, t, :], labels[:, t])
459
+
460
+ # 等价于一次性算全部
461
+ logits = model(input_ids)
462
+ loss = CE(shift_logits, shift_labels)
463
+ ```
464
+
465
+ ### 问题:Exposure Bias
466
+
467
+ ```
468
+ 训练时:
469
+ input: "中国的首都是" → 模型预测 → "北京"
470
+ 实际输入下一时间步: "北京" (真实 token)
471
+
472
+ 推理时:
473
+ input: "中国的首都是" → 模型预测 → "上海" (错误!)
474
+ 实际输入下一时间步: "上海" (自己的预测, 错上加错)
475
+
476
+ 训练分布 ≠ 推理分布 → 累积误差
477
+ ```
478
+
479
+ ### 缓解方法
480
+
481
+ 1. **Scheduled Sampling**:推理时以一定概率用模型自己的预测替换真实 token(本仓库未实现,但面试常考)
482
+ 2. **强化学习**(RLHF阶段):直接在自生成序列上优化
483
+ 3. **Beam Search**:推理时维护候选路径,减少单步错误的累积影响
484
+
485
+ ---
486
+
487
+ ## Q12. 实际训练资源估算
488
+
489
+ ### 30M 模型 SFT 成本
490
+
491
+ | 项目 | 估算 |
492
+ |---|---|
493
+ | 参数量 | ~30M(hidden_size=512, L=8) |
494
+ | 总步数 | `ceil(60000/16) × 2 = 7500` |
495
+ | 每步时间 | ~250ms (RTX 4060) |
496
+ | 总时间 | `7500 × 0.25 ≈ 31 分钟` |
497
+ | 峰值显存 | ~4-5 GB(bf16, bs=16, seq=768) |
498
+ | 权重大小 | ~60 MB(fp16 保存) |
499
+
500
+ ### 大数据全量 SFT 估算(实际生产)
501
+
502
+ | 数据量 | batch_size | 步数 | 每步时间 | 总时间 |
503
+ |---|---|---|---|---|
504
+ | 10K | 16 | 1250 | ~250ms | ~5 分钟 |
505
+ | 60K | 16 | 7500 | ~250ms | ~31 分钟 |
506
+ | 500K | 16 | 62500 | ~250ms | ~4.3 小时 |
507
+
508
+ > 面试点:如何加速 SFT 训练?→ 增大 batch_size(需更大显存或多卡)→ 减少步数;使用梯度累积补偿显存不足;使用 DeepSpeed ZeRO 节省显存
509
+
510
+ ---
511
+
512
+ ## 面试高频题汇总
513
+
514
+ ### 基础
515
+
516
+ 1. **SFT 和 Pretrain 训练的核心区别?** → 数据格式(纯文本 vs 对话)、loss 计算(全 token vs assistant only)、学习率(5e-4 vs 1e-5)、权重初始化
517
+ 2. **为什么 label 要置 -100?** → `CrossEntropyLoss(ignore_index=-100)` 忽略该位置损失,只计算 assistant 段
518
+ 3. **Teacher Forcing 是什么?** → 训练时每步输入真实 token 而非模型预测
519
+ 4. **灾难性遗忘怎么避免?** → 小 lr、少轮数、混合预训练数据、LoRA 增量微调
520
+
521
+ ### 进阶
522
+
523
+ 5. **Exposure Bias 是什么?** → 训练(teacher forcing)和推理(自回归)的输入分布不一致导致的误差累积
524
+ 6. **SFT 后为什么需要 RLHF?** → SFT 只是模仿,RLHF 从偏好信号中学习"什么更好"
525
+ 7. **AdamW 比 Adam 好在哪?** → weight decay 与自适应学习率解耦,更有效的正则化
526
+ 8. **weight tying 在微调时有用吗?** → 有用,嵌入层和输出头共享权重能提升泛化和收敛
527
+ 9. **strict=False 的潜在风险?** → 部分参数随机初始化而不报错,需要手动验证加载结果
528
+ 10. **如何处理超长对话?** → 截断(丢失信息)、滑动窗口(窗口训练)、压缩(使用长上下文模型)
docs/interview/pretrain.md ADDED
@@ -0,0 +1,520 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 面试:预训练(Pretrain)深度
2
+
3
+ > 本仓库 `src/trainers/lm/pretrain.py`,对应 config:`configs/lm/lm_pretrain.yaml`
4
+
5
+ ## 0. 整体流程
6
+
7
+ ```
8
+ 数据集 (pretrain_t2t_mini.jsonl)
9
+ │ 1,270,238 条纯文本
10
+
11
+ PretrainDataset
12
+ ├── 读取 jsonl,每条取 text 字段
13
+ ├── tokenizer.encode() → input_ids
14
+ ├── labels = input_ids 克隆(完整序列监督)
15
+ └── 截断 / 填充到 max_seq_len
16
+
17
+
18
+ SkipBatchSampler → DataLoader
19
+ │ batch_size 个样本 / 组
20
+ │ SkipBatchSampler 支持断点续训跳过前 N 步
21
+
22
+ Forward: model(input_ids, labels)
23
+ ├── TokenEmbed → N×Block → RMSNorm → lm_head
24
+ └── CE Loss(input=logits[..., :-1, :], target=labels[..., 1:])
25
+
26
+
27
+ Backward → 梯度累积 → 梯度裁剪 → Optimizer.step()
28
+
29
+
30
+ Cosine LR 调度 → 定期保存 checkpoint
31
+ ```
32
+
33
+ ---
34
+
35
+ ## Q1. PretrainDataset 怎么处理数据?
36
+
37
+ ### 代码(`src/dataset/pretrain.py`)
38
+
39
+ ```python
40
+ class PretrainDataset(Dataset):
41
+ def __init__(self, data_path, tokenizer, max_length=340):
42
+ self.data = self.load_data(data_path)
43
+ self.tokenizer = tokenizer
44
+ self.max_length = max_length
45
+
46
+ def load_data(self, data_path):
47
+ with open(data_path, 'r') as f:
48
+ return [json.loads(line)['text'] for line in f]
49
+
50
+ def __getitem__(self, index):
51
+ text = self.data[index]
52
+ input_ids = self.tokenizer.encode(text, max_length=self.max_length,
53
+ truncation=True, padding='max_length',
54
+ return_tensors='pt')[0]
55
+ labels = input_ids.clone()
56
+ return input_ids, labels
57
+ ```
58
+
59
+ ### 与 SFT 数据集的关键区别
60
+
61
+ | | PretrainDataset | SFTDataset |
62
+ |---|---|---|
63
+ | 数据格式 | 纯文本 | 多轮对话 JSON |
64
+ | label 构造 | `labels = input_ids.clone()`(全监督) | prompt 段置 `-100`,只标 assistant |
65
+ | 损失计算 | 所有 token 都参与 | 只有 assistant 回复参与 |
66
+ | 训练目标 | 下一个 token 预测 | 指令跟随回复生成 |
67
+
68
+ > 面试点:预训练时所有 token 都计算损失 vs SFT 只算 assistant 段,为什么?→ 预训练目标是学习语言分布,所有 token 都提供统计信息;SFT 是学习指令跟随能力,只需要拟合回复
69
+
70
+ ---
71
+
72
+ ## Q2. 损失函数具体怎么算?
73
+
74
+ ### Next-Token Prediction + 标签平移
75
+
76
+ ```python
77
+ # src/models/lm/model.py:63-70
78
+ def forward(self, input_ids, labels=None):
79
+ hidden_states = self.model(input_ids)
80
+ logits = self.lm_head(hidden_states) # [B, S, V]
81
+
82
+ if labels is not None:
83
+ shift_logits = logits[..., :-1, :].contiguous() # [B, S-1, V]
84
+ shift_labels = labels[..., 1:].contiguous() # [B, S-1]
85
+ loss = F.cross_entropy(
86
+ shift_logits.view(-1, shift_logits.size(-1)),
87
+ shift_labels.view(-1)
88
+ )
89
+ ```
90
+
91
+ ### 为什么平移?
92
+
93
+ - `logits[t]` 预测的是 `labels[t+1]`(第 t 个位置的输出预测第 t+1 个位置的 token)
94
+ - 不平移的话,模型学到的是恒等映射:`logits[t] ≈ labels[t]`
95
+ - 训练时看似 loss 很低,但生成时完全无法产生新内容
96
+
97
+ > 面试点:平移的本质是什么?→ 将语言模型建模为条件概率 P(xₜ|x₁,...,xₜ₋₁),第 t 个位置的输出对应第 t+1 个 token 的概率分布
98
+
99
+ ### 辅助损失:MoE aux_loss
100
+
101
+ ```python
102
+ loss = res.loss + res.aux_loss
103
+ ```
104
+
105
+ 当 `use_moe=1` 时,MoE 层还会产生一个 auxiliary load-balancing loss,鼓励专家负载均衡。非 MoE 模式下 `aux_loss = 0`。
106
+
107
+ ---
108
+
109
+ ## Q3. 训练超参如何影响模型?
110
+
111
+ ### 核心参数详解
112
+
113
+ | 参数 | 默认值 | 作用 | 调大影响 | 调小影响 |
114
+ |---|---|---|---|---|
115
+ | `hidden_size` | 512 | 每 token 的表示维度 | 容量↑,速度↓ | 容量↓,速度↑ |
116
+ | `num_hidden_layers` | 8 | Transformer 层数 | 深度↑,梯度传播难 | 深度↓,表达弱 |
117
+ | `batch_size` | 32 | 每步样本数 | 梯度稳,显存↑ | 梯度噪,显存↓ |
118
+ | `accumulation_steps` | 8 | 梯度累积步数 | 等效 batch↑,速度不变 | 等效 batch↓ |
119
+ | `max_seq_len` | 340 | 最大序列长度 | 上下文↑,显存↑ | 上下文↓ |
120
+ | `learning_rate` | 5e-4 | 初始学习率 | 收敛快,可能不稳 | 收敛慢,更稳 |
121
+ | `dtype` | bfloat16 | 混合精度类型 | 精度高,速度中 | 精度低,速度快 |
122
+
123
+ ### 等效 Batch Size
124
+
125
+ ```
126
+ 等效 batch = batch_size × accumulation_steps
127
+ ```
128
+
129
+ 本仓库 pretrain 默认 `batch_size=32, accumulation_steps=8` → 等效 batch = 256。
130
+
131
+ ### 梯度累积实现
132
+
133
+ ```python
134
+ # src/trainers/lm/pretrain.py:36-47
135
+ loss = loss / args.accumulation_steps # 归一化
136
+ scaler.scale(loss).backward() # 累积梯度
137
+
138
+ if step % args.accumulation_steps == 0: # 累积够 N 步后
139
+ scaler.unscale_(optimizer)
140
+ clip_grad_norm_(model.parameters(), grad_clip)
141
+ scaler.step(optimizer) # 更新参数
142
+ scaler.update()
143
+ optimizer.zero_grad(set_to_none=True) # 清空梯度
144
+ ```
145
+
146
+ > 面试点:为什么 loss 要除以 accumulation_steps?→ 因为梯度是链式累积的,不归一化的话等效学习率会放大 accumulation_steps 倍。除以步数后,每步梯度相当于独立 batch 梯度的平均,保持学习率语义不变
147
+
148
+ ---
149
+
150
+ ## Q4. 学习率调度策略?
151
+
152
+ ### Cosine Decay
153
+
154
+ ```python
155
+ # src/utils/training.py:82-83
156
+ def get_lr(current_step, total_steps, lr):
157
+ return lr * (0.1 + 0.45 * (1 + math.cos(math.pi * current_step / total_steps)))
158
+ ```
159
+
160
+ ### 调度曲线
161
+
162
+ ```
163
+ lr
164
+
165
+ │ lr * 0.55 ─────────── 余弦下降 ──→ lr * 0.1
166
+ │ (初始) (最终)
167
+ └─────────────────────────────────────→ step
168
+ ```
169
+
170
+ - 初始 LR = `lr * 0.55`(cos 从 0 开始,`1 + cos(0) = 2` → `0.1 + 0.45*2 = 1.0`... 不对)
171
+ 仔细看:`current_step=0` 时 `cos(0)=1` → `0.1 + 0.45*2 = 1.0` → `lr * 1.0 = lr`
172
+ 所以在 step 0 时 LR 正好等于设置的 learning_rate。
173
+ - 最终 LR = `lr * 0.1`(cos(π) = -1 → `0.1 + 0.45*0 = 0.1`)
174
+ - 即 LR 从 `lr` 余弦衰减到 `0.1 * lr`
175
+
176
+ 这种"不下到 0"的调度比标准 cosine 更好,保持模型在训练后期仍有适度更新能力。
177
+
178
+ ---
179
+
180
+ ## Q5. 混合精度训练怎么做?
181
+
182
+ ### AMP (Automatic Mixed Precision)
183
+
184
+ ```python
185
+ # src/trainers/lm/pretrain.py:121-123
186
+ dtype = torch.bfloat16 if args.dtype == "bfloat16" else torch.float16
187
+ autocast_ctx = torch.cuda.amp.autocast(dtype=dtype)
188
+ ```
189
+
190
+ ### bf16 vs fp16
191
+
192
+ | | bf16 | fp16 |
193
+ |---|---|---|
194
+ | 指数位 | 8 位(同 fp32) | 5 位 |
195
+ | 尾数位 | 7 位 | 10 位 |
196
+ | 数值范围 | 同 fp32(~3.4e38) | 有限(~6.5e4) |
197
+ | 是否需要 GradScaler | ❌ 不需要 | ✅ 需要 |
198
+ | 精度 | 低精度(7bit 尾数) | 高精度(10bit 尾数) |
199
+ | 硬件要求 | Ampere+(3090/A100 等) | 几乎所有 GPU |
200
+
201
+ ### 为什么 bf16 不需要 GradScaler?
202
+
203
+ bf16 的指数范围和 fp32 一样,不会发生梯度下溢。fp16 的指数范围只有 5 位,小梯度会直接变 0,需要用 GradScaler 放大梯度再缩小。
204
+
205
+ ### 精度保留技巧
206
+
207
+ ```python
208
+ # src/core/norm.py:7-8
209
+ x = x.float() # RMSNorm 内部转 fp32
210
+ x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
211
+ return (self.weight * x).type_as(self.weight) # 转回 bf16/fp16
212
+ ```
213
+
214
+ 即使训练是 bf16,归一化层内部转 fp32 计算再转回,避免归一化精度损失。
215
+
216
+ ---
217
+
218
+ ## Q6. 训练时显存都花在哪了?
219
+
220
+ ### 显存分布(以 30M 参数模型为例)
221
+
222
+ ```
223
+ 模型参数 (fp32): 120 MB
224
+ 梯度 (fp32): 120 MB
225
+ Adam 状态 (fp32×2): 240 MB
226
+ ─────────────────────────────
227
+ 模型状态总计: 480 MB
228
+
229
+ 激活值 (bf16, 8层): ~3-5 GB ← 大头
230
+ 输入数据: <10 MB
231
+ CUDA 上下文: ~200 MB
232
+ ─────────────────────────────
233
+ 总计: ~4-6 GB(取决于 batch_size 和 seq_len)
234
+ ```
235
+
236
+ ### 为什么激活值占这么多?
237
+
238
+ - 反向传播需要存储每层的中间激活值
239
+ - 每层存储 Q/K/V、注意力输出、MLP 中间结果等
240
+ - 数量级:`O(batch_size × seq_len × hidden_size × num_layers × k)`,k ≈ 15-34
241
+
242
+ ### 没有激活检查点
243
+
244
+ 本仓库**没有**使用 `torch.utils.checkpoint`(梯度检查点)。代价是激活值全存,好处是不需要重计算,训练速度更快。
245
+
246
+ > 面试点:激活检查点 trade-off 是什么?→ 节省显存(存部分激活,反向时重算),增加约 15-20% 计算时间
247
+
248
+ ---
249
+
250
+ ## Q7. 分布式训练是怎么做的?
251
+
252
+ ### DDP (DistributedDataParallel)
253
+
254
+ ```python
255
+ # src/utils/distributed.py
256
+ def init_distributed_mode():
257
+ if int(os.environ.get("RANK", -1)) == -1:
258
+ return 0 # 单卡模式
259
+ dist.init_process_group(backend="nccl") # 多卡模式
260
+ local_rank = int(os.environ["LOCAL_RANK"])
261
+ torch.cuda.set_device(local_rank)
262
+ return local_rank
263
+ ```
264
+
265
+ ### 启动方式
266
+
267
+ ```bash
268
+ # 单卡
269
+ python -m trainers.lm.pretrain --config configs/lm/lm_pretrain.yaml
270
+
271
+ # 多卡(torchrun)
272
+ torchrun --nproc_per_node=4 -m trainers.lm.pretrain --config configs/lm/lm_pretrain.yaml
273
+ ```
274
+
275
+ ### DDP 原理
276
+
277
+ - 每个 GPU 一张完整的模型副本
278
+ - 前向/反向独立计算
279
+ - 反向传播后通过 `allreduce` 同步梯度
280
+ - 每个 GPU 独立执行 optimizer.step()
281
+
282
+ ### 本仓库没有使用
283
+
284
+ - ❌ ZeRO(DeepSpeed)
285
+ - ❌ FSDP
286
+ - ❌ 张量/序列并行
287
+ - ❌ torch.compile(默认关闭,`use_compile: 0`)
288
+
289
+ > 面试点:为什么小模型不用 ZeRO/FSDP?→ 模型仅 30M 参数,单卡就能装下,DDP 的梯度同步开销也很小。ZeRO/FSDP 的通信量更大,对小模型反而可能更慢
290
+
291
+ ---
292
+
293
+ ## Q8. Checkpoint 如何保存和恢复?
294
+
295
+ ### 保存(`train_epoch` 内)
296
+
297
+ ```python
298
+ # src/trainers/lm/pretrain.py:59-67
299
+ if (step % args.save_interval == 0 or step == iters) and is_main_process():
300
+ ckp = f'{args.save_dir}/{args.save_weight}_{hidden_size}.pth'
301
+ torch.save({k: v.half().cpu() for k, v in state_dict.items()}, ckp)
302
+ lm_checkpoint(lm_config, weight=args.save_weight, model=model,
303
+ optimizer=optimizer, epoch=epoch, step=step, ...)
304
+ ```
305
+
306
+ ### 保存两个文件
307
+
308
+ | 文件 | 内容 | 用途 |
309
+ |---|---|---|
310
+ | `pretrain_512.pth` | 模型权重 (fp16) | 推理/下游微调 |
311
+ | `pretrain_512_resume.pth` | 权重 + 优化器 + epoch + step | 断点续训 |
312
+
313
+ ### 恢复训练 `from_resume=1`
314
+
315
+ ```python
316
+ ckp_data = lm_checkpoint(lm_config, weight=args.save_weight, save_dir='../checkpoints')
317
+ if ckp_data:
318
+ model.load_state_dict(ckp_data['model'])
319
+ optimizer.load_state_dict(ckp_data['optimizer'])
320
+ start_epoch = ckp_data['epoch']
321
+ start_step = ckp_data.get('step', 0)
322
+ ```
323
+
324
+ ### 权重初始化 `from_weight`
325
+
326
+ ```yaml
327
+ from_weight: none # 从头训练(随机初始化)
328
+ from_weight: pretrain # 加载 pretrain_512.pth 继续训练
329
+ from_weight: full_sft # 加载 full_sft_512.pth 继续训练
330
+ ```
331
+
332
+ `init_model` 会根据 `from_weight` 在 `save_dir` 下查找对应文件:
333
+
334
+ ```python
335
+ weight_path = f'{save_dir}/{from_weight}_{hidden_size}.pth'
336
+ weights = torch.load(weight_path, map_location=device)
337
+ model.load_state_dict(weights, strict=False)
338
+ ```
339
+
340
+ > 面试点:`strict=False` 意味着什么?→ 允许加载的权重和模型结构部分不匹配(如只加载 encoder 不加载 lm_head),LoRA 等场景常用
341
+
342
+ ---
343
+
344
+ ## Q9. SkipBatchSampler 如何实现断点续训?
345
+
346
+ ```python
347
+ # src/utils/training.py:177-200
348
+ class SkipBatchSampler(Sampler):
349
+ def __init__(self, sampler, batch_size, skip_batches=0):
350
+ self.sampler = sampler # 原始索引(或 DistributedSampler)
351
+ self.batch_size = batch_size
352
+ self.skip_batches = skip_batches # 跳过的步数
353
+
354
+ def __iter__(self):
355
+ batch = []
356
+ skipped = 0
357
+ for idx in self.sampler:
358
+ batch.append(idx)
359
+ if len(batch) == self.batch_size:
360
+ if skipped < self.skip_batches: # 跳过前 N 个 batch
361
+ skipped += 1
362
+ batch = []
363
+ continue
364
+ yield batch
365
+ batch = []
366
+ ```
367
+
368
+ ### 为什么需要跳过?
369
+
370
+ - 断点续训时已经从 checkpoint 恢复了优化器状态
371
+ - 但 DataLoader 从头开始迭代的话,会重复处理之前的数据
372
+ - `SkipBatchSampler` 通过 `skip_batches` 跳过已处理的 batch
373
+
374
+ ### 使用场景
375
+
376
+ ```python
377
+ skip = start_step if (epoch == start_epoch and start_step > 0) else 0
378
+ batch_sampler = SkipBatchSampler(train_sampler or indices, args.batch_size, skip)
379
+ ```
380
+
381
+ ---
382
+
383
+ ## Q10. 随机种子和数据打乱
384
+
385
+ ```python
386
+ # src/trainers/lm/pretrain.py:160
387
+ setup_seed(42 + epoch)
388
+ indices = torch.randperm(len(train_ds)).tolist()
389
+ ```
390
+
391
+ ### 每 epoch 重新打乱
392
+
393
+ - 每个 epoch 用不同的 `indices` 顺序
394
+ - 种子 = `42 + epoch`,保证可复现
395
+ - `DistributedSampler` 模式下,每个 GPU 拿到不同但确定性的分片
396
+
397
+ ### 为什么 seed 要 + epoch?
398
+
399
+ - 每个 epoch 的数据顺序不同
400
+ - 同一 epoch 在不同运行间可复现
401
+ - 分布式下每个 rank 拿到不同的子集
402
+
403
+ ---
404
+
405
+ ## Q11. YAML 配置是如何生效的?
406
+
407
+ ### 配置优先级
408
+
409
+ ```
410
+ CLI 参数 > YAML 默认值 > Python argparse 默认值
411
+ ```
412
+
413
+ ### 实现机制
414
+
415
+ ```python
416
+ # src/utils/training.py:25-38
417
+ def apply_config(parser, default_config=None):
418
+ pre, _ = parser.parse_known_args()
419
+ config_path = getattr(pre, 'config', None) or default_config
420
+ if config_path and os.path.exists(config_path):
421
+ defaults = _load_yaml_config(config_path)
422
+ parser.set_defaults(**defaults) # YAML 值设置为 argparse 默认值
423
+ return parser.parse_args() # CLI 显式传参仍可覆盖
424
+ ```
425
+
426
+ ### YAML 映射
427
+
428
+ ```yaml
429
+ model:
430
+ hidden_size: 512 → args.hidden_size
431
+ train:
432
+ batch_size: 32 → args.batch_size
433
+ paths:
434
+ save_dir: ... → args.save_dir
435
+ ```
436
+
437
+ YAML 的三个 section(model/train/paths)被扁平化为 argparse 参数后注入,之后任何 CLI 传参都可以覆盖。
438
+
439
+ ---
440
+
441
+ ## Q12. 为什么选 SwiGLU 作为激活函数?
442
+
443
+ ### 公式对比
444
+
445
+ | 激活函数 | 公式 | 参数量 |
446
+ |---|---|---|
447
+ | ReLU FFN | `ReLU(xW₁)W₂` | 2×d×d_ff |
448
+ | SwiGLU | `SiLU(xW_g) * (xW_u) * W_d` | 3×d×d_ff |
449
+
450
+ SwiGLU 用三个投影(gate/up/down)替代两个,参数量增加 50%,但同等参数下效果更好。
451
+
452
+ ### 仓库实现
453
+
454
+ ```python
455
+ # src/core/mlp.py:16-17
456
+ def forward(self, x):
457
+ return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
458
+ ```
459
+
460
+ ### intermediate_size 的奇怪计算
461
+
462
+ ```python
463
+ # src/models/lm/config.py:22
464
+ self.intermediate_size = math.ceil(hidden_size * math.pi / 64) * 64
465
+ ```
466
+
467
+ 为���么用 π?这是一个取整技巧:`hidden_size × π / 64` 取整再 ×64,保证 intermediate_size 是 64 的倍数,有利于 GPU 内存对齐和 Tensor Core 加速。
468
+
469
+ 对于 `hidden_size=512`:`intermediate_size = ceil(512 × π / 64) × 64 = ceil(25.13) × 64 = 1664`
470
+
471
+ ---
472
+
473
+ ## Q13. 训练过程中 loss 的变化规律
474
+
475
+ 典型 pretrain loss 曲线:
476
+
477
+ ```
478
+ loss
479
+
480
+ 8.0 │ █
481
+ 7.0 │ ██
482
+ 6.0 │ ███
483
+ 5.0 │ ████
484
+ 4.0 │ █████
485
+ 3.0 │ ██████
486
+ └────────────────────→ step
487
+ ```
488
+
489
+ ### 特征
490
+
491
+ 1. **快速下降期**(前 5-10% 步数):loss 从 ~8.5 降到 ~5.0,模型学到基础词法/语法模式
492
+ 2. **平稳下降期**:loss 稳步下降,学习更复杂的语义/知识
493
+ 3. **接近收敛**:loss 下降变缓,接近理论下界
494
+
495
+ ### 预估最终 loss
496
+
497
+ 对于 vocab_size=6400 的随机初始化模型:
498
+ - 初始 loss ≈ log(6400) ≈ 8.76(均匀分布的交叉熵)
499
+ - 训练后 loss ≈ 2.5-3.5(取决于模型大小和数据量)
500
+ - 理论下限 ≈ 0(完美拟合数据分布,但实际上达不到)
501
+
502
+ ---
503
+
504
+ ## 面试高频题汇总
505
+
506
+ ### 基础
507
+
508
+ 1. **预训练和 SFT 的区别?** → 预训练从零学习语言分布(全 token 监督),SFT 学习指令跟随(只监督回复)
509
+ 2. **为什么用 bf16 而不是 fp16?** → bf16 指数范围同 fp32,无需 GradScaler,训练更稳定
510
+ 3. **梯度累积的作用?** → 显存不足时用计算换显存,等效增大 batch_size
511
+ 4. **Cosine 学习率调度的优缺点?** → 平滑衰减,早期快速学习后期精细调优,但可能过早衰减
512
+
513
+ ### 进阶
514
+
515
+ 5. **如何估计训练时间?** → `总步数 = ceil(样本数 / batch_size) × epochs`,`总时间 = 总步数 × 每步时间`
516
+ 6. **为什么 loss 除以 accumulation_steps?** → 保持梯度期望不变,等效于 "平均 N 个 mini-batch 的梯度"
517
+ 7. **DDP 和 DP 的区别?** → DDP 每个 GPU 独立前反向 + allreduce 梯度,DP 是单进程多线程(GIL 限制性能)
518
+ 8. **weight tying 的作用?** → 共享 embedding 和 lm_head 的权重矩阵,减少参数量(本项目默认开启)
519
+ 9. **GQA 和 MHA 的区别?** → GQA 减少 KV head 数量,降低 KV cache 大小,推理时更省显存
520
+ 10. **Flash Attention 为什么省显存?** → tiling 计算注意力矩阵,不显式存储 O(n²) 的 score 矩阵
docs/training/config-and-cli.md CHANGED
@@ -53,10 +53,11 @@ trainer 用 `LMConfig(**vars(args))`(或 `VLMConfig` / `VAMConfig`)构造模
53
  ## 启动
54
 
55
  ```bash
56
- bash runs/train_lm.sh # 默认 configs/lm.yaml
57
- bash runs/train_vlm.sh --config configs/vlm_moe.yaml
58
- bash runs/train_vam.sh --epochs 10 # 覆盖单字段
59
- bash runs/train_tokenizer.sh # 训练 tokenizer(学习用),保存到 checkpoint/tokenizer/
 
60
  ```
61
 
62
  ## Tokenizer 训练
@@ -70,10 +71,10 @@ bash runs/train_tokenizer.sh # 训练 tokenizer(学习用
70
  常用参数:
71
 
72
  ```bash
73
- bash runs/train_tokenizer.sh --data_path dataset/sft_t2t_mini.jsonl \
74
- --vocab_size 6400 \
75
- --checkpoint_dir ../checkpoint \
76
- --no_eval
77
  ```
78
 
79
  ## 要点(面试)
 
53
  ## 启动
54
 
55
  ```bash
56
+ python -m trainers.lm.full_sft --config configs/lm/lm_full_sft.yaml
57
+ python -m trainers.vlm.full_sft --config configs/vlm/vlm_moe.yaml
58
+ python -m trainers.vam.full_sft --config configs/vam/vam.yaml --epochs 10
59
+ python -m trainers.lm.train_tokenizer --data_path dataset/sft_t2t_mini.jsonl \
60
+ --vocab_size 6400 --no_eval
61
  ```
62
 
63
  ## Tokenizer 训练
 
71
  常用参数:
72
 
73
  ```bash
74
+ python -m trainers.lm.train_tokenizer --data_path dataset/sft_t2t_mini.jsonl \
75
+ --vocab_size 6400 \
76
+ --checkpoint_dir ../checkpoint \
77
+ --no_eval
78
  ```
79
 
80
  ## 要点(面试)
docs/training/index.md CHANGED
@@ -24,7 +24,7 @@
24
 
25
  ## 训练脚本(Trainers)
26
 
27
- 按模态组织的训练入口,每个脚本暴露 `main(default_config=None)`(可由 `python -m trainers.<mod>` 或 `runs/*.sh` 调用)。
28
 
29
  - `trainers/lm/`:pretrain / full_sft / lora / dpo / distillation / ppo / grpo / agent / rollout_engine / train_tokenizer
30
  - `trainers/vlm/`:pretrain / full_sft
 
24
 
25
  ## 训练脚本(Trainers)
26
 
27
+ 按模态组织的训练入口,每个脚本暴露 `main(default_config=None)`(通过 `python -m trainers.<mod>` 调用)。
28
 
29
  - `trainers/lm/`:pretrain / full_sft / lora / dpo / distillation / ppo / grpo / agent / rollout_engine / train_tokenizer
30
  - `trainers/vlm/`:pretrain / full_sft
docs/training/trainers.md CHANGED
@@ -1,6 +1,6 @@
1
  # trainers/trainers.md · 训练脚本
2
 
3
- `trainers/` 按模态分子包,每个脚本暴露 `main(default_config=None)`(可由 `python -m trainers.<mod>` 或 `runs/*.sh` 调用)。
4
 
5
  ## 文本(lm/)
6
 
 
1
  # trainers/trainers.md · 训练脚本
2
 
3
+ `trainers/` 按模态分子包,每个脚本暴露 `main(default_config=None)`(通过 `python -m trainers.<mod>` 调用)。
4
 
5
  ## 文本(lm/)
6
 
runs/train_lm.sh DELETED
@@ -1,6 +0,0 @@
1
- #!/usr/bin/env bash
2
- # LM (纯文本) 训练启动脚本
3
- set -e
4
- ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
5
- cd "$ROOT"
6
- exec uv run python -m trainers.lm.full_sft --config "$ROOT/configs/lm.yaml" "$@"
 
 
 
 
 
 
 
runs/train_tokenizer.sh DELETED
@@ -1,8 +0,0 @@
1
- #!/usr/bin/env bash
2
- # Tokenizer 训练启动脚本(仅供学习和参考,不建议重复训练 tokenizer)
3
- # 用法: bash runs/train_tokenizer.sh [额外参数...]
4
- set -e
5
- ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
6
- cd "$ROOT"
7
- export PYTHONPATH="$ROOT/src:${PYTHONPATH}"
8
- exec uv run python -m trainers.lm.train_tokenizer "$@"
 
 
 
 
 
 
 
 
 
runs/train_vam.sh DELETED
@@ -1,6 +0,0 @@
1
- #!/usr/bin/env bash
2
- # VAM (文本 + 视觉 + 语音 全模态) 训练启动脚本
3
- set -e
4
- ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
5
- cd "$ROOT"
6
- exec uv run python -m trainers.vam.full_sft --config "$ROOT/configs/vam.yaml" "$@"
 
 
 
 
 
 
 
runs/train_vlm.sh DELETED
@@ -1,6 +0,0 @@
1
- #!/usr/bin/env bash
2
- # VLM (文本 + 视觉) 训练启动脚本
3
- set -e
4
- ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
5
- cd "$ROOT"
6
- exec uv run python -m trainers.vlm.full_sft --config "$ROOT/configs/vlm.yaml" "$@"
 
 
 
 
 
 
 
src/trainers/lm/full_sft.py CHANGED
@@ -13,72 +13,12 @@ from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
  from models import LMConfig
15
  from dataset import SFTDataset
16
- from utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
17
  from utils.training import apply_config # noqa: F401
18
 
19
  warnings.filterwarnings('ignore')
20
 
21
 
22
- def train_epoch(epoch, loader, iters, start_step=0, wandb=None):
23
- start_time = time.time()
24
- last_step = start_step
25
- for step, (input_ids, labels) in enumerate(loader, start=start_step + 1):
26
- input_ids = input_ids.to(args.device)
27
- labels = labels.to(args.device)
28
- last_step = step
29
- lr = get_lr(epoch * iters + step, args.epochs * iters, args.learning_rate)
30
- for param_group in optimizer.param_groups:
31
- param_group['lr'] = lr
32
-
33
- with autocast_ctx:
34
- res = model(input_ids, labels=labels)
35
- loss = res.loss + res.aux_loss
36
- loss = loss / args.accumulation_steps
37
-
38
- scaler.scale(loss).backward()
39
-
40
- if step % args.accumulation_steps == 0:
41
- scaler.unscale_(optimizer)
42
- torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
43
-
44
- scaler.step(optimizer)
45
- scaler.update()
46
-
47
- optimizer.zero_grad(set_to_none=True)
48
-
49
- if step % args.log_interval == 0 or step == iters:
50
- spend_time = time.time() - start_time
51
- current_loss = loss.item() * args.accumulation_steps
52
- current_aux_loss = res.aux_loss.item() if res.aux_loss is not None else 0.0
53
- current_logits_loss = current_loss - current_aux_loss
54
- current_lr = optimizer.param_groups[-1]['lr']
55
- eta_min = spend_time / max(step - start_step, 1) * (iters - step) // 60
56
- Logger(f'Epoch:[{epoch + 1}/{args.epochs}]({step}/{iters}), loss: {current_loss:.4f}, logits_loss: {current_logits_loss:.4f}, aux_loss: {current_aux_loss:.4f}, lr: {current_lr:.8f}, epoch_time: {eta_min:.1f}min')
57
- if wandb: wandb.log({"loss": current_loss, "logits_loss": current_logits_loss, "aux_loss": current_aux_loss, "learning_rate": current_lr, "epoch_time": eta_min})
58
-
59
- if (step % args.save_interval == 0 or step == iters) and is_main_process():
60
- model.eval()
61
- moe_suffix = '_moe' if lm_config.use_moe else ''
62
- ckp = f'{args.save_dir}/{args.save_weight}_{lm_config.hidden_size}{moe_suffix}.pth'
63
- raw_model = model.module if isinstance(model, DistributedDataParallel) else model
64
- raw_model = getattr(raw_model, '_orig_mod', raw_model)
65
- state_dict = raw_model.state_dict()
66
- torch.save({k: v.half().cpu() for k, v in state_dict.items()}, ckp)
67
- lm_checkpoint(lm_config, weight=args.save_weight, model=model, optimizer=optimizer,
68
- epoch=epoch, step=step, wandb=wandb, save_dir='../checkpoints', scaler=scaler)
69
- model.train()
70
- del state_dict
71
-
72
- del input_ids, labels, res, loss
73
-
74
- if last_step > start_step and last_step % args.accumulation_steps != 0:
75
- scaler.unscale_(optimizer)
76
- torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
77
- scaler.step(optimizer)
78
- scaler.update()
79
- optimizer.zero_grad(set_to_none=True)
80
-
81
-
82
  def main(default_config=None):
83
  parser = argparse.ArgumentParser(description="MiniMind Full SFT")
84
  parser.add_argument('--config', type=str, default=None, help='YAML 配置路径,其字段作为 argparse 默认值,CLI 显式传参可覆盖')
@@ -100,6 +40,7 @@ def main(default_config=None):
100
  parser.add_argument('--max_seq_len', default=768, type=int, help="训练的最大截断长度(中文1token≈1.5~1.7字符)")
101
  parser.add_argument('--use_moe', default=0, type=int, choices=[0, 1], help="是否使用MoE架构(0=否,1=是)")
102
  parser.add_argument("--data_path", type=str, default="../dataset/sft_t2t_mini.jsonl", help="训练数据路径")
 
103
  parser.add_argument('--from_weight', default='pretrain', type=str, help="基于哪个权重训练,为none则不基于任何权重训练")
104
  parser.add_argument('--from_resume', default=0, type=int, choices=[0, 1], help="是否自动检测&续训(0=否,1=是)")
105
  parser.add_argument("--use_wandb", action="store_true", help="是否使用wandb")
@@ -133,7 +74,7 @@ def main(default_config=None):
133
  wandb.init(project=args.wandb_project, name=wandb_run_name, id=wandb_id, resume=resume)
134
 
135
  # ========== 5. 定义模型、数据、优化器 ==========
136
- model, tokenizer = init_model(lm_config, args.from_weight, save_dir=args.save_dir, tokenizer_dir=args.tokenizer_dir, device=args.device)
137
  train_ds = SFTDataset(args.data_path, tokenizer, max_length=args.max_seq_len)
138
  train_sampler = DistributedSampler(train_ds) if dist.is_initialized() else None
139
  scaler = torch.cuda.amp.GradScaler(enabled=(args.dtype == 'float16'))
@@ -155,7 +96,67 @@ def main(default_config=None):
155
  if dist.is_initialized():
156
  model = DistributedDataParallel(model, device_ids=[local_rank])
157
 
158
- # ========== 8. 开始训练 ==========
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
159
  for epoch in range(start_epoch, args.epochs):
160
  train_sampler and train_sampler.set_epoch(epoch)
161
  setup_seed(42 + epoch); indices = torch.randperm(len(train_ds)).tolist()
@@ -168,7 +169,7 @@ def main(default_config=None):
168
  else:
169
  train_epoch(epoch, loader, len(loader), 0, wandb)
170
 
171
- # ========== 9. 清理分布进程 ==========
172
  if dist.is_initialized():
173
  dist.barrier()
174
  dist.destroy_process_group()
 
13
  from torch.utils.data import DataLoader, DistributedSampler
14
  from models import LMConfig
15
  from dataset import SFTDataset
16
+ from utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler, init_logger
17
  from utils.training import apply_config # noqa: F401
18
 
19
  warnings.filterwarnings('ignore')
20
 
21
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22
  def main(default_config=None):
23
  parser = argparse.ArgumentParser(description="MiniMind Full SFT")
24
  parser.add_argument('--config', type=str, default=None, help='YAML 配置路径,其字段作为 argparse 默认值,CLI 显式传参可覆盖')
 
40
  parser.add_argument('--max_seq_len', default=768, type=int, help="训练的最大截断长度(中文1token≈1.5~1.7字符)")
41
  parser.add_argument('--use_moe', default=0, type=int, choices=[0, 1], help="是否使用MoE架构(0=否,1=是)")
42
  parser.add_argument("--data_path", type=str, default="../dataset/sft_t2t_mini.jsonl", help="训练数据路径")
43
+ parser.add_argument('--model_dir', default=None, type=str, help="预训练权重所在目录(默认同save_dir)")
44
  parser.add_argument('--from_weight', default='pretrain', type=str, help="基于哪个权重训练,为none则不基于任何权重训练")
45
  parser.add_argument('--from_resume', default=0, type=int, choices=[0, 1], help="是否自动检测&续训(0=否,1=是)")
46
  parser.add_argument("--use_wandb", action="store_true", help="是否使用wandb")
 
74
  wandb.init(project=args.wandb_project, name=wandb_run_name, id=wandb_id, resume=resume)
75
 
76
  # ========== 5. 定义模型、数据、优化器 ==========
77
+ model, tokenizer = init_model(lm_config, args.from_weight, save_dir=args.save_dir, tokenizer_dir=args.tokenizer_dir, device=args.device, model_dir=args.model_dir)
78
  train_ds = SFTDataset(args.data_path, tokenizer, max_length=args.max_seq_len)
79
  train_sampler = DistributedSampler(train_ds) if dist.is_initialized() else None
80
  scaler = torch.cuda.amp.GradScaler(enabled=(args.dtype == 'float16'))
 
96
  if dist.is_initialized():
97
  model = DistributedDataParallel(model, device_ids=[local_rank])
98
 
99
+ # ========== 8. 训练函数 ==========
100
+ def train_epoch(epoch, loader, iters, start_step=0, wandb=None):
101
+ start_time = time.time()
102
+ last_step = start_step
103
+ for step, (input_ids, labels) in enumerate(loader, start=start_step + 1):
104
+ input_ids = input_ids.to(args.device)
105
+ labels = labels.to(args.device)
106
+ last_step = step
107
+ lr = get_lr(epoch * iters + step, args.epochs * iters, args.learning_rate)
108
+ for param_group in optimizer.param_groups:
109
+ param_group['lr'] = lr
110
+
111
+ with autocast_ctx:
112
+ res = model(input_ids, labels=labels)
113
+ loss = res.loss + res.aux_loss
114
+ loss = loss / args.accumulation_steps
115
+
116
+ scaler.scale(loss).backward()
117
+
118
+ if step % args.accumulation_steps == 0:
119
+ scaler.unscale_(optimizer)
120
+ torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
121
+
122
+ scaler.step(optimizer)
123
+ scaler.update()
124
+
125
+ optimizer.zero_grad(set_to_none=True)
126
+
127
+ if step % args.log_interval == 0 or step == iters:
128
+ spend_time = time.time() - start_time
129
+ current_loss = loss.item() * args.accumulation_steps
130
+ current_aux_loss = res.aux_loss.item() if res.aux_loss is not None else 0.0
131
+ current_logits_loss = current_loss - current_aux_loss
132
+ current_lr = optimizer.param_groups[-1]['lr']
133
+ eta_min = spend_time / max(step - start_step, 1) * (iters - step) // 60
134
+ Logger(f'Epoch:[{epoch + 1}/{args.epochs}]({step}/{iters}), loss: {current_loss:.4f}, logits_loss: {current_logits_loss:.4f}, aux_loss: {current_aux_loss:.4f}, lr: {current_lr:.8f}, epoch_time: {eta_min:.1f}min')
135
+ if wandb: wandb.log({"loss": current_loss, "logits_loss": current_logits_loss, "aux_loss": current_aux_loss, "learning_rate": current_lr, "epoch_time": eta_min})
136
+
137
+ if (step % args.save_interval == 0 or step == iters) and is_main_process():
138
+ model.eval()
139
+ moe_suffix = '_moe' if lm_config.use_moe else ''
140
+ ckp = f'{args.save_dir}/{args.save_weight}_{lm_config.hidden_size}{moe_suffix}.pth'
141
+ raw_model = model.module if isinstance(model, DistributedDataParallel) else model
142
+ raw_model = getattr(raw_model, '_orig_mod', raw_model)
143
+ state_dict = raw_model.state_dict()
144
+ torch.save({k: v.half().cpu() for k, v in state_dict.items()}, ckp)
145
+ lm_checkpoint(lm_config, weight=args.save_weight, model=model, optimizer=optimizer,
146
+ epoch=epoch, step=step, wandb=wandb, save_dir='../checkpoints', scaler=scaler)
147
+ model.train()
148
+ del state_dict
149
+
150
+ del input_ids, labels, res, loss
151
+
152
+ if last_step > start_step and last_step % args.accumulation_steps != 0:
153
+ scaler.unscale_(optimizer)
154
+ torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
155
+ scaler.step(optimizer)
156
+ scaler.update()
157
+ optimizer.zero_grad(set_to_none=True)
158
+
159
+ # ========== 9. 开始训练 ==========
160
  for epoch in range(start_epoch, args.epochs):
161
  train_sampler and train_sampler.set_epoch(epoch)
162
  setup_seed(42 + epoch); indices = torch.randperm(len(train_ds)).tolist()
 
169
  else:
170
  train_epoch(epoch, loader, len(loader), 0, wandb)
171
 
172
+ # ========== 10. 清理分布进程 ==========
173
  if dist.is_initialized():
174
  dist.barrier()
175
  dist.destroy_process_group()
src/trainers/lm/pretrain.py CHANGED
@@ -13,7 +13,7 @@ from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
  from models import LMConfig
15
  from dataset import PretrainDataset
16
- from utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
17
  from utils.training import apply_config # noqa: F401
18
 
19
  warnings.filterwarnings('ignore')
 
13
  from torch.utils.data import DataLoader, DistributedSampler
14
  from models import LMConfig
15
  from dataset import PretrainDataset
16
+ from utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler, init_logger
17
  from utils.training import apply_config # noqa: F401
18
 
19
  warnings.filterwarnings('ignore')
src/utils/training.py CHANGED
@@ -158,14 +158,15 @@ def lm_checkpoint(lm_config, weight='full_sft', model=None, optimizer=None, epoc
158
  return None
159
 
160
 
161
- def init_model(lm_config, from_weight='pretrain', save_dir='../checkpoint', tokenizer_dir=None, device='cuda'):
162
  tokenizer_dir = tokenizer_dir or os.path.join(save_dir, 'tokenizer')
163
  tokenizer = AutoTokenizer.from_pretrained(tokenizer_dir)
164
  model = LMForCausalLM(lm_config)
165
 
166
  if from_weight != 'none':
167
  moe_suffix = '_moe' if lm_config.use_moe else ''
168
- weight_path = f'{save_dir}/{from_weight}_{lm_config.hidden_size}{moe_suffix}.pth'
 
169
  weights = torch.load(weight_path, map_location=device)
170
  model.load_state_dict(weights, strict=False)
171
 
 
158
  return None
159
 
160
 
161
+ def init_model(lm_config, from_weight='pretrain', save_dir='../checkpoint', tokenizer_dir=None, device='cuda', model_dir=None):
162
  tokenizer_dir = tokenizer_dir or os.path.join(save_dir, 'tokenizer')
163
  tokenizer = AutoTokenizer.from_pretrained(tokenizer_dir)
164
  model = LMForCausalLM(lm_config)
165
 
166
  if from_weight != 'none':
167
  moe_suffix = '_moe' if lm_config.use_moe else ''
168
+ weight_dir = model_dir or save_dir
169
+ weight_path = f'{weight_dir}/{from_weight}_{lm_config.hidden_size}{moe_suffix}.pth'
170
  weights = torch.load(weight_path, map_location=device)
171
  model.load_state_dict(weights, strict=False)
172