Files
ars-opd-rebuild/ars_opd/configs.py
T
iomgaa a0faec0df7 层1: 修复远程 sanity OOM——per_device batch 8→2、累积 2→8(全局 64 不变)
根因: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 <noreply@anthropic.com>
2026-07-18 09:21:38 -04:00

156 lines
7.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""实验配置(当前只含层 1 的 SFTConfig,后续层在此文件追加各自的 dataclass)。
设计约定(对应 CLAUDE.md §2"配置显式化"):
- 所有实验参数必须是这里某个 dataclass 的字段;代码里出现魔法数字/路径即违规。
- 密钥(API key 等)不进配置类,走 `.env`(见 teacher.py)。
- 机器相关路径(数据集、输出目录)不给默认值,强制调用方显式传入——
防止参考实现里 `/fsx` 硬编码那类"在别人机器上必炸"的坑。
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class SFTConfig:
"""层 1 SFT 基线的全部实验参数。
论文锚点:§3.1 式(1) 的标准交叉熵 SFT;但按 §5.1 的基线定义,
训练数据是 teacher rollout(离线蒸馏),不是人写答案——所以有
`teacher_completions_path` 字段:DAPO 是 prompt-only 数据集,
解答一律来自 teacher 生成的缓存文件。
frozen=True:配置一旦构造即只读。训练中途被悄悄改掉的配置是最难
排查的 bug 来源之一;要换参数就构造一个新实例,留下明确的代码痕迹。
"""
# ---- 机器相关路径(无默认值,必须显式传入)----
dataset_path: str
"""DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。"""
output_dir: str
"""checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。"""
# ---- 数据 ----
teacher_completions_path: str | None = None
"""teacher 解答缓存(teacher.py 生成的 JSONL)。None 表示数据集自带
assistant 轮次;若数据实际是 prompt-only 又没给此路径,data.py 会显式报错,
不做静默兜底。"""
dataset_split: str = "train"
subset_size: int | None = 1000
"""随机抽取的子集大小(控制 teacher API 成本,roadmap 定为 ~1k);None = 全量。"""
# ---- 序列双预算 ----
# 非显然约束:prompt 与 completion 必须各有独立预算。若只用一个 max_length
# 从右截断,超长解答会把 prompt 挤空,模型在"没有题目"的样本上学习解答
# ——这是参考实现 collator(trainer:267-292) 的头号正确性卖点,此处继承。
max_length: int = 4096
"""prompt + completion 的总 token 预算。"""
max_prompt_length: int = 1024
"""prompt 单独预算;completion 实际预算 = max_length - len(截断后 prompt)。"""
enable_thinking: bool = False
"""Qwen3 chat 模板的思考开关。False 时模板注入空 `<think>\\n\\n</think>`。
非显然约束:此开关改变渲染后的 prompt 文本,从而改变 prompt/completion
的 token 边界——训练与推理必须取同一值,否则掩码整体错位。"""
# ---- 优化 ----
# 差异标注:论文 §5.1 的蒸馏训练用 lr=1e-6,参考实现 SFT 默认 2e-5
# (train_distillation.py:73)。SFT 有真实 token 监督、信号密集,从参考实现取
# 2e-5;层 5 的蒸馏配置再回到论文的 1e-6。
learning_rate: float = 2e-5
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 = 8
"""全局 batch = 2(per_device) × 4(卡) × 8(累积) = 64,与参考实现注释的
训练规模(trainer 配置注释"global batch 64")对齐。"""
num_train_epochs: int = 1
max_steps: int = -1
""">0 时覆盖 num_train_epochs,只跑这么多步——远程 50 步 sanity 用;-1 = 按 epoch。"""
lr_scheduler_type: str = "linear"
warmup_ratio: float = 0.0
gradient_checkpointing: bool = False
"""0.6B 学生显存富余,不开(省 ~40% 显存、慢 ~30%)。注意:FSDP 下此开关
是 no-op,真正的开关是 FSDP_ACTIVATION_CHECKPOINTING 环境变量(见
scripts/ 训练脚本头部的前置块,docs/02 §2.6)。"""
bf16: bool = True
seed: int = 42
# ---- 日志与保存 ----
logging_steps: int = 1
save_steps: int = 100
save_total_limit: int = 2
report_to: str = "none"
""""none" 或 "wandb"。默认 none:本地调试不该悄悄往外发数据,远程脚本显式开。"""
def __post_init__(self) -> None:
"""构造即校验:配置错误必须在训练开始前炸,而不是跑到第一个超长样本才炸。"""
if self.max_prompt_length >= self.max_length:
raise ValueError(
f"max_prompt_length({self.max_prompt_length}) 必须小于 "
f"max_length({self.max_length}),否则 completion 预算为零,"
f"所有样本的 labels 将全为 -100,loss 恒为 0 且无报错——静默空训练。"
)
if self.learning_rate <= 0:
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
if self.subset_size is not None and self.subset_size <= 0:
raise ValueError(
f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}"
)
if self.max_steps == 0 or self.max_steps < -1:
raise ValueError(
f"max_steps 只接受 -1(按 epoch)或正整数,收到 {self.max_steps}"
)
@dataclass(frozen=True)
class TeacherGenConfig:
"""teacher 批量生成(层 1 能力)的采样与执行参数。
连接信息(API 地址/密钥/模型名)不在这里——那是部署环境的事实,走 `.env`
teacher.py 读取);这里只放"换一组值就是换一个实验"的采样参数。
"""
temperature: float = 1.0
top_p: float = 0.95
"""MiniMax M 系官方推荐采样参数:temperature=1.0, top_p=0.95。"""
max_tokens: int = 16384
"""teacher 单条回复的 token 上限。这是上限不是目标——按实际生成量计费,
放大它不增加正常解答的成本,只给最难的题留出写完的空间(8192 时 59 条实测
截断 2 条)。非显然约束:M3 的思考段也计入此额度,设太小会把解答挤没。"""
strip_think: bool = True
"""剥离 content 开头的 <think>...</think> 思考段。SFT 的监督目标是最终
解答;student 以 enable_thinking=False 训练,学思考段会与模板约定矛盾。"""
concurrency: int = 16
"""并发请求数(线程池大小)。上限看网关的承受力,报 429 就调小。"""
max_retries: int = 3
"""单请求的网络级重试次数(openai 客户端内建指数退避)。"""
system_prompt: str | None = None
"""None = 不加 system 轮(DAPO 题面自带作答指令,不需要额外指挥)。"""
def __post_init__(self) -> None:
if self.max_tokens <= 0:
raise ValueError(f"max_tokens 必须为正,收到 {self.max_tokens}")
if self.concurrency < 1:
raise ValueError(f"concurrency 必须 ≥1,收到 {self.concurrency}")
if self.temperature < 0:
raise ValueError(f"temperature 必须 ≥0,收到 {self.temperature}")