From 58d75cc56ac8140b48cc8cfb1deeb673d8e4dec9 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 04:40:45 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=821/T1:=20SFTConfig=20dataclass=20?= =?UTF-8?q?=E4=B8=8E=E6=9E=84=E9=80=A0=E6=A0=A1=E9=AA=8C=EF=BC=88docs/02?= =?UTF-8?q?=20=C2=A74=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- ars_opd/__init__.py | 4 ++ ars_opd/configs.py | 113 ++++++++++++++++++++++++++++++++++++++++++ tests/test_configs.py | 49 ++++++++++++++++++ 3 files changed, 166 insertions(+) create mode 100644 ars_opd/__init__.py create mode 100644 ars_opd/configs.py create mode 100644 tests/test_configs.py diff --git a/ars_opd/__init__.py b/ars_opd/__init__.py new file mode 100644 index 0000000..15950f7 --- /dev/null +++ b/ars_opd/__init__.py @@ -0,0 +1,4 @@ +"""ars_opd:OmniOPD (arXiv:2606.01476v2) 的分层重构实现。 + +模块与论文的对应关系见 CLAUDE.md §3(单一事实源),此处不重复。 +""" diff --git a/ars_opd/configs.py b/ars_opd/configs.py new file mode 100644 index 0000000..81dc7fd --- /dev/null +++ b/ars_opd/configs.py @@ -0,0 +1,113 @@ +"""实验配置(当前只含层 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 时模板注入空 `\\n\\n`。 + 非显然约束:此开关改变渲染后的 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}" + ) diff --git a/tests/test_configs.py b/tests/test_configs.py new file mode 100644 index 0000000..277f3f0 --- /dev/null +++ b/tests/test_configs.py @@ -0,0 +1,49 @@ +"""SFTConfig 的构造校验测试(层 1 / T1)。 + +只测"配置错误必须在构造时炸"这一条约定;参数语义本身没有逻辑可测。 +""" + +import dataclasses + +import pytest + +from ars_opd.configs import SFTConfig + + +def make(**overrides): + """最小合法配置;单测只关心被覆盖的那个字段。""" + base = dict(dataset_path="dummy.parquet", output_dir="/tmp/dummy") + base.update(overrides) + return SFTConfig(**base) + + +def test_合法配置可构造(): + cfg = make() + assert cfg.max_length > cfg.max_prompt_length + + +def test_prompt预算吞掉总预算时报错(): + # 这是最危险的静默失败:completion 预算为 0 → labels 全 -100 → loss 恒 0 + with pytest.raises(ValueError, match="max_prompt_length"): + make(max_prompt_length=4096, max_length=4096) + + +def test_非法学习率报错(): + with pytest.raises(ValueError, match="learning_rate"): + make(learning_rate=0.0) + + +def test_非法子集大小报错(): + with pytest.raises(ValueError, match="subset_size"): + make(subset_size=0) + + +def test_非法max_steps报错(): + with pytest.raises(ValueError, match="max_steps"): + make(max_steps=0) + + +def test_配置冻结不可变(): + cfg = make() + with pytest.raises(dataclasses.FrozenInstanceError): + cfg.learning_rate = 1e-3