Files
ars-opd-rebuild/ars_opd/configs.py
T
iomgaa 58d75cc56a 层1/T1: SFTConfig dataclass 与构造校验(docs/02 §4)
- ars_opd/configs.py: 冻结 dataclass,机器路径无默认强制显式传入;
  双预算/enable_thinking/lr 差异均按 §7 规范标注
- __post_init__ 构造即校验,防"completion 预算为零→loss 恒 0"静默空训练
- tests/test_configs.py: 6 个校验测试

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 04:40:45 -04:00

114 lines
5.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 = 8
gradient_accumulation_steps: int = 2
"""全局 batch = 8(per_device) × 4(卡) × 2(累积) = 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}"
)