reorganize configs into subdirs lm/vlm/vam, remove runs/, add interview docs for pretrain & full_sft
Browse files- README.md +75 -19
- configs/{lm.yaml → lm/lm_full_sft.yaml} +2 -1
- configs/lm/lm_full_sft_mini.yaml +27 -0
- configs/{lm_moe.yaml → lm/lm_full_sft_moe.yaml} +0 -0
- configs/lm/lm_pretrain.yaml +26 -0
- configs/lm/lm_pretrain_mini.yaml +26 -0
- configs/lm/lm_pretrain_moe.yaml +31 -0
- configs/{vam.yaml → vam/vam.yaml} +0 -0
- configs/{vam_moe.yaml → vam/vam_moe.yaml} +0 -0
- configs/{vlm.yaml → vlm/vlm.yaml} +0 -0
- configs/{vlm_moe.yaml → vlm/vlm_moe.yaml} +0 -0
- docs/architecture/overview.md +3 -4
- docs/interview/full_sft.md +528 -0
- docs/interview/pretrain.md +520 -0
- docs/training/config-and-cli.md +9 -8
- docs/training/index.md +1 -1
- docs/training/trainers.md +1 -1
- runs/train_lm.sh +0 -6
- runs/train_tokenizer.sh +0 -8
- runs/train_vam.sh +0 -6
- runs/train_vlm.sh +0 -6
- src/trainers/lm/full_sft.py +65 -64
- src/trainers/lm/pretrain.py +1 -1
- src/utils/training.py +3 -2
README.md
CHANGED
|
@@ -67,18 +67,15 @@ src/
|
|
| 67 |
│ └── checkpoint.py # checkpoint 读写辅助
|
| 68 |
├── serve/ # 实时语音会话(SileroVAD / RealtimeSession)
|
| 69 |
configs/
|
| 70 |
-
├──
|
| 71 |
-
├──
|
|
|
|
|
|
|
| 72 |
├── vlm.yaml # 视觉多模态训练配置
|
| 73 |
-
├──
|
| 74 |
-
├──
|
| 75 |
-
├──
|
| 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 |
-
|
| 106 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 107 |
|
| 108 |
```bash
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
--no_eval
|
| 113 |
|
| 114 |
-
|
| 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/
|
|
|
|
| 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 |
-
|
| 45 |
-
|
| 46 |
-
|
| 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 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
|
|
|
| 60 |
```
|
| 61 |
|
| 62 |
## Tokenizer 训练
|
|
@@ -70,10 +71,10 @@ bash runs/train_tokenizer.sh # 训练 tokenizer(学习用
|
|
| 70 |
常用参数:
|
| 71 |
|
| 72 |
```bash
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 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)`(
|
| 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)`(
|
| 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 |
-
# ==========
|
| 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 |
-
|
|
|
|
| 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 |
|