From a0faec0df7b5d5e88e5479c9fdb371f0f23f08f3 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 09:21:38 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=821:=20=E4=BF=AE=E5=A4=8D=E8=BF=9C?= =?UTF-8?q?=E7=A8=8B=20sanity=20OOM=E2=80=94=E2=80=94per=5Fdevice=20batch?= =?UTF-8?q?=208=E2=86=922=E3=80=81=E7=B4=AF=E7=A7=AF=202=E2=86=928?= =?UTF-8?q?=EF=BC=88=E5=85=A8=E5=B1=80=2064=20=E4=B8=8D=E5=8F=98=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 根因:150k 大词表下显存大头是 (B,T,V) logits 链(fp32 ~20G@B=8)与逐层激活, 均正比于 B 而与 0.6B 参数量无关。docs/02 §2.6 旧显存估算勘误入档; train_sft.sh 加 expandable_segments 防碎片。 Co-Authored-By: Claude Fable 5 --- .vscode/settings.json | 4 ++++ ars_opd/configs.py | 9 ++++++--- docs/02-sft-baseline.md | 2 +- scripts/train_sft.py | 4 ++-- scripts/train_sft.sh | 1 + 5 files changed, 14 insertions(+), 6 deletions(-) create mode 100644 .vscode/settings.json diff --git a/.vscode/settings.json b/.vscode/settings.json new file mode 100644 index 0000000..4b5a294 --- /dev/null +++ b/.vscode/settings.json @@ -0,0 +1,4 @@ +{ + "python-envs.defaultEnvManager": "ms-python.python:conda", + "python-envs.defaultPackageManager": "ms-python.python:conda" +} \ No newline at end of file diff --git a/ars_opd/configs.py b/ars_opd/configs.py index 67a8d48..0c74c0c 100644 --- a/ars_opd/configs.py +++ b/ars_opd/configs.py @@ -64,10 +64,13 @@ class SFTConfig: # 2e-5;层 5 的蒸馏配置再回到论文的 1e-6。 learning_rate: float = 2e-5 - per_device_train_batch_size: int = 8 + per_device_train_batch_size: int = 2 + """非显然约束:别看 0.6B 小就调大它——显存大头是 (B,T,V) 的 logits 链 + (fp32 一份 ~20G@B=8)与逐层激活,都正比于 B 而与参数量无关;B=8 实测 + 爆 80G 卡(2026-07-18 远程 sanity)。""" - gradient_accumulation_steps: int = 2 - """全局 batch = 8(per_device) × 4(卡) × 2(累积) = 64,与参考实现注释的 + gradient_accumulation_steps: int = 8 + """全局 batch = 2(per_device) × 4(卡) × 8(累积) = 64,与参考实现注释的 训练规模(trainer 配置注释"global batch 64")对齐。""" num_train_epochs: int = 1 diff --git a/docs/02-sft-baseline.md b/docs/02-sft-baseline.md index d447840..a030cf6 100644 --- a/docs/02-sft-baseline.md +++ b/docs/02-sft-baseline.md @@ -58,7 +58,7 @@ F.cross_entropy(..., ignore_index=-100) 参考实现用 torchrun + FSDP,有个著名 trick:`FSDP_ACTIVATION_CHECKPOINTING` 环境变量必须在 import accelerate/transformers **之前**设置(train_dist:15-21),且 `TrainingArguments.gradient_checkpointing` 在 FSDP 下是 no-op(train_dist:390-393)。 -**我们的决策(2026-07-18 讨论定)**:0.6B 乃至 1.7B 学生 4×A800 都用 **DDP**(显存账:1.7B 全套 ~27G/卡,80G 舒适);FSDP 触发点 = **换 4B 学生或序列 >8K**。但"import 前设环境变量"这个坑**今天就消除**:T5 训练脚本头部从第一天起内置 `os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", ...)` 前置块——DDP 下无害 no-op,未来启用 FSDP 时只改 TrainingArguments 字段,核心模块零改动。启用后必须 `nvidia-smi` 实测显存验证生效(分布式配置"设置了但静默无效"是常态,不可信配置)。另:层 2-5 坚持 DDP 还有调试纯度考量——on-policy 生成/ref model 搬运与 FSDP 的交界是参考实现最毛的地方,先保证"出错必是算法错"。 +**我们的决策(2026-07-18 讨论定)**:0.6B 乃至 1.7B 学生 4×A800 都用 **DDP**;FSDP 触发点 = **换 4B 学生或序列 >8K**。⚠️ 显存账勘误(2026-07-18 sanity 实爆教训):"1.7B 全套 ~27G/卡"的旧估算只算了参数系(参数+梯度+Adam),漏了两个与参数量无关、正比于 batch 的大头:(B,T,V) logits 链(fp32 一份即 B×T×151936×4 字节,B=8/T=4096 时 ~20G,cross_entropy 内部 log_softmax 再来一份)与逐层激活(~30-40G@B=8)。150k 大词表下**显存瓶颈是 B×T,不是模型大小**;对策 = 压 per_device batch 用梯度累积补(B=2×4卡×8累积=64 不变)。但"import 前设环境变量"这个坑**今天就消除**:T5 训练脚本头部从第一天起内置 `os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", ...)` 前置块——DDP 下无害 no-op,未来启用 FSDP 时只改 TrainingArguments 字段,核心模块零改动。启用后必须 `nvidia-smi` 实测显存验证生效(分布式配置"设置了但静默无效"是常态,不可信配置)。另:层 2-5 坚持 DDP 还有调试纯度考量——on-policy 生成/ref model 搬运与 FSDP 的交界是参考实现最毛的地方,先保证"出错必是算法错"。 ## 3. 保留 / 替代 / 删除清单(CLAUDE.md §6.2 规定动作) diff --git a/scripts/train_sft.py b/scripts/train_sft.py index ea23ace..874ac95 100644 --- a/scripts/train_sft.py +++ b/scripts/train_sft.py @@ -35,8 +35,8 @@ FULL = SFTConfig( max_prompt_length=1024, enable_thinking=False, learning_rate=2e-5, - per_device_train_batch_size=8, - gradient_accumulation_steps=2, # 全局 batch = 8 × 4 卡 × 2 = 64 + per_device_train_batch_size=2, # B=8 曾爆 80G:大头是 (B,T,V) logits 链与激活,见 SFTConfig 注释 + gradient_accumulation_steps=8, # 全局 batch = 2 × 4 卡 × 8 = 64 num_train_epochs=1, max_steps=-1, lr_scheduler_type="linear", diff --git a/scripts/train_sft.sh b/scripts/train_sft.sh index 9757c3d..6f5e97d 100644 --- a/scripts/train_sft.sh +++ b/scripts/train_sft.sh @@ -19,6 +19,7 @@ MODE=${1:-full} export CUDA_VISIBLE_DEVICES=$GPUS export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5) +export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长序列 batch 易碎片化,按需扩段 export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com} export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满