diff --git a/examples/on_policy_distillation/mopd/run-mopd-qwen3-vl-2b-8xgpu-colocate.sh b/examples/on_policy_distillation/mopd/run-mopd-qwen3-vl-2b-8xgpu-colocate.sh index e114d6756..eabb000c7 100755 --- a/examples/on_policy_distillation/mopd/run-mopd-qwen3-vl-2b-8xgpu-colocate.sh +++ b/examples/on_policy_distillation/mopd/run-mopd-qwen3-vl-2b-8xgpu-colocate.sh @@ -164,6 +164,7 @@ SGLANG_ARGS=( --sglang-mem-fraction-static 0.8 --sglang-load-format dummy --sglang-enable-weights-cpu-backup + --sglang-disable-cuda-graph ) RESOURCE_JSON="{\"actor\": [1, ${ACTOR_GPUS}], \"rollout\": [1, ${ROLLOUT_GPUS}], \"teacher\": [1, ${TEACHER_GPUS}]}" diff --git a/examples/on_policy_distillation/mopd/run-mopd-qwen35-35ba3b-16xgpu-colocate.sh b/examples/on_policy_distillation/mopd/run-mopd-qwen35-35ba3b-16xgpu-colocate.sh index 5e9ff2f53..f06f116cf 100644 --- a/examples/on_policy_distillation/mopd/run-mopd-qwen35-35ba3b-16xgpu-colocate.sh +++ b/examples/on_policy_distillation/mopd/run-mopd-qwen35-35ba3b-16xgpu-colocate.sh @@ -177,6 +177,7 @@ SGLANG_ARGS=( --sglang-max-running-requests 128 --sglang-load-format dummy --sglang-enable-weights-cpu-backup + --sglang-disable-cuda-graph ) PARTIAL_ROLLOUT_ARGS=( diff --git a/examples/on_policy_distillation/mopd/run-mopd-qwen35-9b-8xgpu-colocate.sh b/examples/on_policy_distillation/mopd/run-mopd-qwen35-9b-8xgpu-colocate.sh index 660a9ff02..1595239b7 100755 --- a/examples/on_policy_distillation/mopd/run-mopd-qwen35-9b-8xgpu-colocate.sh +++ b/examples/on_policy_distillation/mopd/run-mopd-qwen35-9b-8xgpu-colocate.sh @@ -164,6 +164,7 @@ SGLANG_ARGS=( --sglang-mem-fraction-static 0.7 --sglang-load-format dummy --sglang-enable-weights-cpu-backup + --sglang-disable-cuda-graph ) RESOURCE_JSON="{\"actor\": [1, ${ACTOR_GPUS}], \"rollout\": [1, ${ROLLOUT_GPUS}], \"teacher\": [1, ${TEACHER_GPUS}]}" diff --git a/examples/on_policy_distillation/sdpo/README.md b/examples/on_policy_distillation/sdpo/README.md new file mode 100644 index 000000000..f7ac40d9c --- /dev/null +++ b/examples/on_policy_distillation/sdpo/README.md @@ -0,0 +1,416 @@ +# Relax-SDPO 示例 + +本目录提供 Relax-SDPO 的文本训练示例。SDPO(Self-Distillation from Preference +Optimization)是一种让模型学会「照做成功回答」的训练方法:先让同一个问题生成多个回答, +再用 reward 找出其中的成功回答,把成功回答作为示范注入 teacher prompt,最后让学生模型 +通过 on-policy distillation loss 学习修正自己的回答。 + +这里的所有 launcher 都帮你把整套流程跑通:rollout → reward → feedback → managed +SGLang teacher → `student_topk + JSD` 蒸馏训练,你只需要准备数据和修改环境配置。 + +> **运行前请先确认当前 pod 上没有其他 Ray 任务。** `env.sh` 会执行 `ray stop`,因此 +> source 它会停止当前 pod 上的 Ray 进程。训练机上如果有其他任务正在运行,不要直接执行 +> 这些 launcher。 + +## 这个示例包含什么 + +- 六个开箱即用的四卡 colocate 训练脚本(SciKnowEval × 4、ToolUse、ToolAlpaca) +- 数据转换工具 `prepare_data.py`:把参考数据转成 Relax 的 `prompt`/`label`/`metadata` + JSONL schema +- 规则 reward 与 SDPO feedback 实现(`reward.py` 与 `relax.utils.opd.sdpo.feedback`) +- 由 Relax 自动管理的 SGLang teacher,无需手动部署 + +## 前置条件 + +开始之前请确认以下内容都已就绪: + +- 一台可用 **4×GPU** 的机器 / pod,且当前没有其他 Ray 任务在运行 +- Relax worktree,以及配套的 Relax-SDPO Python 环境(见下方 `env.sh` 的 `RELAX_VENV`) +- Qwen3-8B checkpoint(student 与 teacher 默认共用同一个) +- SDPO 参考数据(SciKnowEval / ToolUse / ToolAlpaca,或你自己的文本数据) +- SGLang 已应用 per-position token-id patch(launcher 会设置 + `RELAX_OPD_PER_POS_TOKEN_IDS=1`),详见 + [通用 OPD 文档中的 SGLang Patch 说明](../README.md#sglang-patch) + +## 快速开始 + +```bash +# 1. 进入项目根目录,确认环境可用(无残留 Ray 任务) +cd +ray status # 或检查当前 pod 的 tmux,确保没有其他任务 +nvidia-smi # 确认四张 GPU 空闲 + +# 2. 准备数据:以 SciKnowEval Chemistry 为例 +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset sciknoweval \ + --input /datasets/sciknoweval/chemistry/train.json \ + --domain chemistry \ + --source-split train \ + --output /SDPO/sciknoweval/chemistry/train.jsonl + +# 3. 修改 examples/on_policy_distillation/sdpo/env.sh: +# 更新 RELAX_VENV / MEGATRON / STUDENT_MODEL_PATH / TEACHER_MODEL_PATH / SDPO_DATA_ROOT + +# 4. 启动训练 +bash examples/on_policy_distillation/sdpo/run-sciknoweval-chemistry-4xgpu-colocate.sh + +# 5. 观察日志与指标(TensorBoard 曲线、训练 loss / reward / teacher 状态等) +``` + +首次验证建议先跑小规模 smoke:`prepare_data.py` 加 `--max-rows 2` 只转两条数据,并临时调小 +launcher 的 `--num-rollout` / `--rollout-batch-size` 后启动。 + +## 训练流程 + +每个 launcher 的一次训练迭代大致经过以下阶段: + +```text +prompt-data + │ + ├── 学生 rollout:同一问题生成多个 response,组成一个 group + │ + ├── custom reward:计算 score,并记录 feedback + │ + ├── SDPO feedback:构造每个样本的 teacher prompt + │ ├── SciKnowEval:共享同 group/UID 内的成功回答,并加入当前样本反馈 + │ └── ToolUse/Code:同样共享同 group/UID 内的成功回答,并加入当前样本反馈 + │ + ├── managed SGLang teacher:对动态 teacher prompt 计算 top-K log-probability + │ + └── student_topk + JSD loss:更新学生模型 +``` + +当前脚本使用以下关键配置: + +| 配置 | 当前值 | 说明 | +| -------------------------------------------- | -------------- | ------------------------------------------------------------------------------ | +| `--use-opd` | 开启 | 启用 on-policy distillation | +| `--opd-type` | `sglang` | teacher 由 Relax 管理的 SGLang 服务提供 log-probability | +| `--opd-token-selection` | `student_topk` | 在学生 rollout 的 top-K token 集合上计算 SDPO 信号 | +| `--opd-log-prob-top-k` | `16` | 每个位置收集 16 个 token 的 log-probability | +| `--opd-kl-type` | `jsd` | 使用 JSD 形式的 token-level distillation criterion | +| `--opd-norm-mode` | `tail` | 保留 top-K 之外的 tail probability mass | +| `--opd-loss-coef` | `1.0` | 将 distillation signal 作为 loss 注入训练 | +| `--opd-kl-coef` | `0.0` | 不使用 advantage 形式的 OPD KL | +| `--opd-disable-rl-reward` | 开启 | 不把基础 RL outcome reward 注入 actor 优化;custom reward 仍用于 SDPO feedback | +| `--group-rm` | 开启 | 让同一个 prompt 的多个 rollout 进入同一 reward group | +| `--use-rollout-logprobs` | 开启 | 复用学生 rollout 阶段的 log-probability 数据 | +| `--colocate` | 开启 | 在 rollout、teacher 和 actor 之间切换共享 GPU 资源 | +| `--teacher-sglang-enable-weights-cpu-backup` | 开启 | colocate sleep/wake 时把 teacher 权重备份到 CPU,避免唤醒后权重丢失 | + +`student_topk` 模式需要 SGLang 支持按位置返回 token ID。launcher 会设置 +`RELAX_OPD_PER_POS_TOKEN_IDS=1`;运行环境还必须安装对应的 SGLang source patch。详见 +[通用 OPD 文档中的 SGLang Patch 说明](../README.md#sglang-patch)。 + +## Launcher 与资源布局 + +所有当前 launcher 都是单机四卡 colocate 配置: + +```text +四张 GPU 的 colocate resource pool +├── actor :4 GPU(TP=4),训练阶段使用整个 pool +├── rollout :3 GPU,rollout 阶段使用 +└── teacher :1 GPU,rollout 阶段使用 managed SGLang teacher +``` + +脚本中的资源配置为: + +```json +{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]} +``` + +六个 launcher 的训练参数完全一致:`--num-rollout 5000`、`--rollout-batch-size 32`、 +`--n-samples-per-prompt 8`、`--global-batch-size 256`、`--eval-interval 5`、 +`--n-samples-per-eval-prompt 16`。 + +| 脚本 | 数据入口 | eval 入口 | Feedback 类 | +| -------------------------------------------------------------------------------------------- | ----------------------------------- | ---------------------------------- | -------------------------- | +| [`run-sciknoweval-biology-4xgpu-colocate.sh`](run-sciknoweval-biology-4xgpu-colocate.sh) | `sciknoweval/biology/train.jsonl` | `sciknoweval/biology/eval.jsonl` | `GoldenAnswerSDPOFeedback` | +| [`run-sciknoweval-chemistry-4xgpu-colocate.sh`](run-sciknoweval-chemistry-4xgpu-colocate.sh) | `sciknoweval/chemistry/train.jsonl` | `sciknoweval/chemistry/eval.jsonl` | `GoldenAnswerSDPOFeedback` | +| [`run-sciknoweval-physics-4xgpu-colocate.sh`](run-sciknoweval-physics-4xgpu-colocate.sh) | `sciknoweval/physics/train.jsonl` | `sciknoweval/physics/eval.jsonl` | `GoldenAnswerSDPOFeedback` | +| [`run-sciknoweval-material-4xgpu-colocate.sh`](run-sciknoweval-material-4xgpu-colocate.sh) | `sciknoweval/material/train.jsonl` | `sciknoweval/material/eval.jsonl` | `GoldenAnswerSDPOFeedback` | +| [`run-tooluse-4xgpu-colocate.sh`](run-tooluse-4xgpu-colocate.sh) | `tooluse/train.jsonl` | `tooluse/eval.jsonl` | `GoldenAnswerSDPOFeedback` | +| [`run-toolalpaca-4xgpu-colocate.sh`](run-toolalpaca-4xgpu-colocate.sh) | `toolalpaca/train.jsonl` | `toolalpaca/eval.jsonl` | `GoldenAnswerSDPOFeedback` | + +当前脚本没有独立的公共 SDPO launcher,每个数据入口都显式指定了自己的 feedback 类。 +student rollout 统一为 `--rollout-max-response-len 8192`;teacher 请求超时由各 launcher 的 +`--opd-teacher-timeout-s` 控制(launcher 显式设置为 600 s),需要时可自行调整。 + +## 文件结构 + +| 文件 | 用途 | +| ------------------------------------ | ------------------------------------------------------------------------------- | +| [`env.sh`](env.sh) | 设置项目根目录、Python/Megatron 环境、模型路径和数据根目录,并停止当前 Ray 进程 | +| [`prepare_data.py`](prepare_data.py) | 将参考数据转换为 Relax 的 `prompt`/`label`/`metadata` JSONL schema | +| [`reward.py`](reward.py) | 提供 SciKnowEval、ToolUse 和 ToolAlpaca 的 rule-based reward | +| `run-*-4xgpu-colocate.sh` | 按数据集启动四卡 colocate SDPO 训练 | + +## 环境准备 + +### 模型 + +当前脚本通过 `scripts/models/qwen3-8B.sh` 加载 Qwen3-8B 的学生模型配置, +并默认使用同一 Qwen3-8B checkpoint 作为 student 和 teacher: + +```text +/Qwen3-8B # student +/Qwen3-8B # teacher +``` + +也可以使用不同的文本 teacher checkpoint,但必须确保 checkpoint 能被当前 managed +SGLang teacher 和 Qwen3-8B 训练配置正确加载。当前 SDPO prompt-routing 路径只支持文本 +输入,不支持多模态字段。 + +### `env.sh` + +launcher 会自行回到项目根目录,并在内部 source `examples/on_policy_distillation/sdpo/env.sh`。 +请先根据当前机器修改其中的环境路径和 checkpoint 路径: + +| 变量 | 用途 | +| -------------------- | ------------------------------------------------------------ | +| `RELAX_VENV` | Relax-SDPO Python 虚拟环境 | +| `RELAX_PYTHON` | 训练入口使用的 Python | +| `MEGATRON` | 与当前 Relax worktree 配套的 Megatron 和 Python package 路径 | +| `PYTHONPATH` | 由项目根目录和 `MEGATRON` 路径组成 | +| `STUDENT_MODEL_PATH` | 学生模型 HF checkpoint | +| `TEACHER_MODEL_PATH` | managed SGLang teacher 的 HF checkpoint | +| `SDPO_DATA_ROOT` | 准备好的 SDPO JSONL 数据根目录 | + +`env.sh` 当前对上述变量使用固定默认值,而不是 `${VAR:-default}` 形式。因此,在命令行 +预先 export `STUDENT_MODEL_PATH` 或 `TEACHER_MODEL_PATH` 会被 `env.sh` 中的赋值覆盖;如果 +需要更换模型或运行环境,应直接修改 `env.sh`,或维护一份本地 launcher/environment 副本。 + +## 数据准备 + +从 Relax 项目根目录执行 `python3 -m examples.on_policy_distillation.sdpo.prepare_data`。 +输入可以是 JSON、JSONL 或 Parquet,输出统一为 JSONL。 + +### 输出 schema + +每一行至少包含 `prompt`、`label` 和 `metadata`: + +```json +{ + "prompt": "question and optional choices", + "label": "gold answer", + "metadata": { + "data_source": "sciknoweval", + "source_split": "train", + "domain": "Chemistry", + "task_type": "mcq" + } +} +``` + +`metadata.data_source` 是 reward 路由键;`metadata` 中的 `answer_key`、`golden_answer` 等 +字段由 reward 使用。不要在 launcher 中把 `--metadata-key metadata` 改成其他字段,除非 +同时修改数据输出 schema 和 reward 逻辑。 + +### SciKnowEval + +数据通常按 domain 保存。以下命令以 Chemistry 为例;Physics、 +Biology 和 Materials 只需要替换 domain、输入路径和输出路径: + +```bash +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset sciknoweval \ + --input /datasets/sciknoweval/chemistry/train.json \ + --domain chemistry \ + --source-split train \ + --output /SDPO/sciknoweval/chemistry/train.jsonl +``` + +也可以处理已经整理成扁平 schema 的参考数据。对于原始 SciKnowEval 格式,转换器会只保留 +L3 样本,并将 `material` 规范化为 `Materials`: + +```bash +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset sciknoweval \ + --input /chemistry/train.json \ + --source-split train \ + --output /SDPO/sciknoweval/chemistry/train.jsonl +``` + +测试 split 可以用同样的命令生成,只需将输入和 `--source-split` 改为 `test`。当前训练 +launcher 默认每 5 个训练迭代使用 `/eval.jsonl` 做一次周期性评测 +(`--eval-prompt-data`)。 + +如果只有 train 源数据、没有独立的 test 集,可以在生成时按比例留出 eval 集: +`--eval-ratio` 会把该比例的规范化行写到 `.parent/eval.jsonl`,其余写 +到 `--output`(train)。launcher 通过 `--eval-prompt-data` 指向这个 `eval.jsonl` +即可周期性评测: + +```bash +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset sciknoweval \ + --input /datasets/sciknoweval/biology/train.json \ + --domain biology \ + --source-split train \ + --eval-ratio 0.1 \ + --seed 42 \ + --output /SDPO/sciknoweval/biology/train.jsonl +# 生成 /SDPO/sciknoweval/biology/train.jsonl + eval.jsonl +``` + +### ToolUse + +如果使用参考 SDPO 仓库中的工具调用数据: + +```bash +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset tooluse \ + --input /datasets/tooluse/train.json \ + --source-split train \ + --output /SDPO/tooluse/train.jsonl +``` + +### ToolAlpaca + +ToolAlpaca 输入通常是 Parquet: + +```bash +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset toolalpaca \ + --input /data/train-00000-of-00001.parquet \ + --source-split train \ + --output /SDPO/toolalpaca/train.jsonl +``` + +读取 Parquet 需要当前 Python 环境安装 `pyarrow`。 + +## 启动训练 + +### 使用 `SDPO_DATA_ROOT` 默认路径 + +如果 `env.sh` 中的 `SDPO_DATA_ROOT` 已经包含以下目录之一,可以直接启动: + +```text +$SDPO_DATA_ROOT/ +├── sciknoweval/ +│ ├── biology/train.jsonl +│ ├── chemistry/train.jsonl +│ ├── material/train.jsonl +│ └── physics/train.jsonl +├── toolalpaca/train.jsonl +└── tooluse/train.jsonl +``` + +例如: + +```bash +bash examples/on_policy_distillation/sdpo/run-sciknoweval-chemistry-4xgpu-colocate.sh +``` + +### 使用自定义数据路径 + +launcher 会优先读取 `DATA_PATH` 环境变量,未设置时才回退到 `SDPO_DATA_ROOT` 下的默认目录, +因此也可以直接指定你自己的数据文件: + +```bash +DATA_PATH=/path/to/my/train.jsonl \ +bash examples/on_policy_distillation/sdpo/run-sciknoweval-chemistry-4xgpu-colocate.sh +``` + +ToolAlpaca 和 ToolUse 的启动方式相同,只需替换 launcher: + +```bash +DATA_PATH=/SDPO/toolalpaca/train.jsonl \ +bash examples/on_policy_distillation/sdpo/run-toolalpaca-4xgpu-colocate.sh + +DATA_PATH=/SDPO/tooluse/train.jsonl \ +bash examples/on_policy_distillation/sdpo/run-tooluse-4xgpu-colocate.sh +``` + +训练开始前,Relax 会启动 managed SGLang teacher,并在 actor、rollout 和 teacher 之间按 +colocate 配置切换 GPU。launcher 默认每 5 个训练迭代在 `/eval.jsonl` 上评测一次, +并在第一个训练迭代前先跑一次 eval 作为基线。 + +## Reward 与 Feedback + +### SciKnowEval + +`reward.py` 要求回答包含 `...` 标签并从中提取选项字母;缺失 `` +标签即视为格式错误(score=0,即使回答里出现了正确选项)。普通选择题会和 +`metadata.answer_key` 比较,true/false 任务会进行归一化比较。 + +三类数据(SciKnowEval、ToolUse/ToolAlpaca、Code)按同一套决策矩阵决定是否进蒸馏、 +注入什么(按优先级);前两者共用 `GoldenAnswerSDPOFeedback`(静态 golden-answer 文本 +任务),Code 预留 `CodeSDPOFeedback` 占位(reward 未接入,调用即报错): + +1. 自身成功(`score >= success_reward_threshold`,默认 1.0,可经 `--opd-feedback-kwargs` 调整)→ 注入同 `group_index`/`metadata.uid` 内成功 peer 的正确解, + 无 peer 时用自己的解,进蒸馏; +2. 失败但有成功 peer → 只注入 peer 的正确解(丢弃当前样本的 feedback),进蒸馏; +3. 失败、无 peer、且是格式/截断错误 → 注入格式/截断反馈文字,进蒸馏; +4. 失败、无 peer、普通算错 → 不注入、不进蒸馏。 + +没有 `group_index` 时使用 `metadata.uid` 隔离;不同 group/UID 之间不会共享回答。 + +失败样本的反馈文字只有两种(截断优先于格式错误),且都不泄露正确答案: + +```text +Your response was truncated because it exceeded the maximum length. +Your answer had the wrong format. The solution must be given in the format: X. +``` + +普通算错不产生任何反馈文字,因此这类样本只能靠同题的 peer 正确解学习。 + +### ToolUse 与 ToolAlpaca + +模型回答需要包含: + +```text +Action: +Action Input: +``` + +reward 分别检查 tool action 和 JSON 参数:缺失 `Action/Action Input` 格式 → 格式反馈; +回答被截断 → 截断反馈(优先于格式);action/input 不匹配视为普通算错 → 无反馈(不泄露 +gold)。进蒸馏的决策与 SciKnowEval 相同:成功/有成功 peer 时注入 peer 正确解;失败且 +无 peer 时只有格式/截断错误才注入反馈文字。 + +## 常见问题 + +### 提示 `Set STUDENT_MODEL_PATH` 或 `Set SDPO_DATA_ROOT` + +检查 `env.sh` 中的模型、Python/Megatron 和 `SDPO_DATA_ROOT` 配置。若使用自定义数据, +可以直接设置 `DATA_PATH`(见上方「使用自定义数据路径」),这样 launcher 不需要依赖 +`SDPO_DATA_ROOT` 对应的默认目录。 + +### 提示 `No rows matched` + +检查 `--dataset` 是否与输入格式匹配。SciKnowEval 原始格式还需要有效的 L3 domain; +ToolAlpaca 输入必须包含 `golden_answer`;ToolUse 输入必须包含参考格式中的 `prompt` 和 +`answer`。 + +### teacher 请求超时或显存不足 + +检查 teacher/rollout 是否确实各分配一张 GPU,并确认 `--colocate` +和 `--teacher-sglang-enable-weights-cpu-backup` 没有被删除。 +所有 launcher 的默认 rollout 规模较大(5000 prompts / global-batch 256);首次验证建议 +先生成小规模 smoke 数据再训练: + +```bash +python3 -m examples.on_policy_distillation.sdpo.prepare_data \ + --dataset sciknoweval \ + --input /datasets/sciknoweval/chemistry/train.json \ + --domain chemistry \ + --source-split train \ + --max-rows 2 \ + --output /SDPO/sciknoweval/chemistry/train-smoke.jsonl +``` + +必要时应在对应 launcher 中调整 rollout 数量、response 长度或 batch 配置。 + +### Top-K log-probability 不可用 + +确认 SGLang 已应用 [通用 OPD 文档中的 per-position token-id patch](../README.md#sglang-patch), +并保留 `RELAX_OPD_PER_POS_TOKEN_IDS=1`。当前 SDPO 路径不能退回到 +`student_sampled`,因为 SDPO prompt routing 只支持 `student_topk`。 + +## 参考 + +- [On-Policy Distillation 通用说明](../README.md) +- [通用 OPD 的 token selection、loss 和 SGLang 配置](../README.md#token-selection-modes) +- [lasgroup/SDPO](https://github.com/lasgroup/SDPO) +- [SciKnowEval](https://github.com/HICAI-ZJU/SciKnowEval) +- [ToolAlpaca](https://huggingface.co/datasets/Ahren09/ToolAlpaca) diff --git a/examples/on_policy_distillation/sdpo/__init__.py b/examples/on_policy_distillation/sdpo/__init__.py new file mode 100644 index 000000000..9c43ccdf1 --- /dev/null +++ b/examples/on_policy_distillation/sdpo/__init__.py @@ -0,0 +1,3 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""Minimal static-teacher SDPO examples for Relax.""" diff --git a/examples/on_policy_distillation/sdpo/prepare_data.py b/examples/on_policy_distillation/sdpo/prepare_data.py new file mode 100644 index 000000000..5ba148b16 --- /dev/null +++ b/examples/on_policy_distillation/sdpo/prepare_data.py @@ -0,0 +1,261 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""Prepare the reference SDPO SciKnowEval and ToolAlpaca data for Relax.""" + +from __future__ import annotations + +import argparse +import json +import logging +import random +from pathlib import Path +from typing import Any, Iterable + + +logger = logging.getLogger(__name__) +TARGET_DOMAINS = frozenset({"Chemistry", "Physics", "Biology", "Materials"}) + + +def _read_jsonl(path: Path) -> list[dict[str, Any]]: + with path.open(encoding="utf-8") as handle: + return [json.loads(line) for line in handle if line.strip()] + + +def _read_rows(path: Path) -> list[dict[str, Any]]: + if path.suffix.lower() in {".jsonl", ".json"}: + if path.suffix.lower() == ".jsonl": + return _read_jsonl(path) + text = path.read_text(encoding="utf-8").strip() + try: + value = json.loads(text) + except json.JSONDecodeError: + return _read_jsonl(path) + return value if isinstance(value, list) else [value] + if path.suffix.lower() == ".parquet": + import pyarrow.parquet as parquet + + return parquet.read_table(path).to_pylist() + raise ValueError(f"Unsupported input format: {path}") + + +def _canonical_domain(value: Any) -> str: + normalized = str(value or "").strip().casefold() + if normalized == "material": + return "Materials" + return normalized.capitalize() + + +def _json_text(value: Any) -> str: + if isinstance(value, str): + return value + return json.dumps(value, ensure_ascii=False, sort_keys=True) + + +def _normalize_sciknoweval_row( + row: dict[str, Any], + *, + source_split: str, + domain: str | None, +) -> dict[str, Any] | None: + if row.get("dataset") == "sciknoweval" and isinstance(row.get("prompt"), str) and "answer" in row: + normalized_domain = _canonical_domain(domain) + if normalized_domain not in {"Chemistry", "Physics", "Biology", "Materials"}: + return None + prompt = str(row["prompt"]).strip() + system = str(row.get("system") or "").strip() + if system: + prompt = f"{system}\n\n{prompt}" + answer = row.get("answer", "") + metadata = { + "data_source": "sciknoweval", + "source_split": source_split, + "domain": normalized_domain, + "task_type": str(row.get("kind", "mcq")), + "answer_key": answer, + "source_index": row.get("idx"), + } + return {"prompt": prompt, "label": _json_text(answer), "metadata": metadata} + + details = row.get("details") or {} + if not isinstance(details, dict) or str(details.get("level", "")).upper() != "L3": + return None + + source_domain = _canonical_domain(row.get("domain")) + if source_domain not in TARGET_DOMAINS: + return None + + choices = row.get("choices") or {} + choice_lines = [ + f"{label}: {text}" for label, text in zip(choices.get("label") or [], choices.get("text") or [], strict=False) + ] + prompt_value = row.get("prompt", {}) + prompt_default = prompt_value.get("default", "") if isinstance(prompt_value, dict) else prompt_value + question = str(row.get("question") or prompt_default).strip() + prompt = question + if choice_lines: + prompt = f"{question}\n\n" + "\n".join(choice_lines) + prompt += "\n\nReason carefully and provide the final answer." + + normalized_domain = source_domain + answer = row.get("answerKey") or row.get("answer", "") + metadata = { + "data_source": "sciknoweval", + "source_split": source_split, + "domain": normalized_domain, + "task_type": str(row.get("type", "unknown")), + "answer_key": answer, + } + return {"prompt": prompt, "label": _json_text(answer), "metadata": metadata} + + +def _normalize_tool_row(row: dict[str, Any], *, source_split: str, dataset: str) -> dict[str, Any] | None: + if row.get("dataset") == "tooluse" and isinstance(row.get("prompt"), str) and "answer" in row: + answer = row.get("answer", "") + try: + golden_answer = json.loads(answer) if isinstance(answer, str) else answer + except json.JSONDecodeError: + golden_answer = answer + prompt = str(row.get("prompt", "")).strip() + metadata = { + "data_source": "tooluse", + "source_split": source_split, + "task_type": str(row.get("kind", "tooluse")), + "golden_answer": golden_answer, + "source_index": row.get("idx"), + } + return {"prompt": prompt, "label": _json_text(answer), "metadata": metadata} + + if dataset != "toolalpaca" or "golden_answer" not in row: + return None + + name = str(row.get("name", "")).strip() + description = str(row.get("description", "")).strip() + documentation = str(row.get("nl_documentation", "")).strip() + instruction = str(row.get("instruction", row.get("prompt", ""))).strip() + prompt = ( + "You are given an API specification and a user request. Select the correct tool and " + "emit the tool call using exactly:\n" + "Action: \nAction Input: \n\n" + f"Tool name: {name}\n" + f"Tool description: {description}\n" + f"Tool documentation:\n{documentation}\n\n" + f"User request:\n{instruction}" + ) + golden_answer = row.get("golden_answer") or [] + metadata = { + "data_source": "toolalpaca", + "source_split": source_split, + "task_type": "tool_call", + "golden_answer": golden_answer, + } + return {"prompt": prompt, "label": _json_text(golden_answer), "metadata": metadata} + + +def _normalize_relax_row(row: dict[str, Any], *, source_split: str) -> dict[str, Any] | None: + if not isinstance(row.get("prompt"), str) or "label" not in row: + return None + metadata = row.get("metadata") + if not isinstance(metadata, dict): + metadata = {} + metadata = dict(metadata) + metadata.setdefault("source_split", source_split) + return {"prompt": row["prompt"], "label": row["label"], "metadata": metadata} + + +def normalize_rows( + dataset: str, + rows: Iterable[dict[str, Any]], + *, + source_split: str, + domain: str | None = None, +) -> list[dict[str, Any]]: + """Convert one supported source schema into Relax's prompt-data schema.""" + if dataset not in {"sciknoweval", "toolalpaca", "tooluse"}: + raise ValueError(f"Unsupported SDPO dataset {dataset!r}") + if source_split not in {"train", "test"}: + raise ValueError(f"Unsupported source split {source_split!r}; expected 'train' or 'test'") + + normalized_rows = [] + for row in rows: + normalized = _normalize_relax_row(row, source_split=source_split) + if normalized is None: + normalized = ( + _normalize_sciknoweval_row(row, source_split=source_split, domain=domain) + if dataset == "sciknoweval" + else _normalize_tool_row(row, source_split=source_split, dataset=dataset) + ) + if normalized is not None: + normalized_rows.append(normalized) + return normalized_rows + + +def _write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + for row in rows: + handle.write(json.dumps(row, ensure_ascii=False) + "\n") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dataset", choices=("sciknoweval", "toolalpaca", "tooluse"), required=True) + parser.add_argument("--input", required=True, type=Path) + parser.add_argument("--output", required=True, type=Path) + parser.add_argument("--source-split", choices=("train", "test"), required=True) + parser.add_argument( + "--domain", + default=None, + help="SciKnowEval domain for the reference flat format; defaults to the input parent directory name.", + ) + parser.add_argument("--max-rows", type=int, default=None, help="Optionally limit output rows for a smoke run.") + parser.add_argument( + "--eval-ratio", + type=float, + default=0.0, + help=( + "Fraction of normalized rows to hold out as an eval set. When >0, the held-out rows are " + "written to /eval.jsonl and the rest to --output (train). " + "Useful for a train/test split when only a single train source is available." + ), + ) + parser.add_argument("--seed", type=int, default=42, help="Seed for the eval/validation split.") + args = parser.parse_args() + + if args.max_rows is not None and args.max_rows < 0: + parser.error("--max-rows must be non-negative") + if not 0.0 <= args.eval_ratio < 1.0: + parser.error("--eval-ratio must be in [0, 1)") + + rows = _read_rows(args.input) + domain = args.domain or args.input.parent.name + normalized = normalize_rows(args.dataset, rows, source_split=args.source_split, domain=domain) + if args.max_rows is not None: + normalized = normalized[: args.max_rows] + if not normalized: + raise ValueError(f"No rows matched dataset={args.dataset!r} from input {args.input}") + + if args.eval_ratio > 0.0: + n_eval = int(round(len(normalized) * args.eval_ratio)) + if n_eval == 0: + raise ValueError( + f"--eval-ratio {args.eval_ratio} with {len(normalized)} rows yields 0 eval rows; " + "raise the ratio or add more input rows." + ) + rng = random.Random(args.seed) + indices = list(range(len(normalized))) + rng.shuffle(indices) + eval_indices = set(indices[:n_eval]) + train_rows = [r for i, r in enumerate(normalized) if i not in eval_indices] + eval_rows = [r for i, r in enumerate(normalized) if i in eval_indices] + _write_jsonl(args.output, train_rows) + eval_path = args.output.with_name("eval.jsonl") + _write_jsonl(eval_path, eval_rows) + logger.info( + f"Split into train ({len(train_rows)} rows) and eval ({len(eval_rows)} rows); eval written to {eval_path}" + ) + else: + _write_jsonl(args.output, normalized) + + +if __name__ == "__main__": + main() diff --git a/examples/on_policy_distillation/sdpo/reward.py b/examples/on_policy_distillation/sdpo/reward.py new file mode 100644 index 000000000..2e70f754c --- /dev/null +++ b/examples/on_policy_distillation/sdpo/reward.py @@ -0,0 +1,162 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""Rule-based rewards and compact feedback for the SDPO examples.""" + +from __future__ import annotations + +import json +import re +from typing import Any + + +def _metadata(sample: Any) -> dict[str, Any]: + value = getattr(sample, "metadata", None) + return value if isinstance(value, dict) else {} + + +def _was_truncated(sample: Any) -> bool: + status = getattr(sample, "status", None) + if status is None: + return False + return str(getattr(status, "value", status)).casefold() == "truncated" + + +def _normalize(value: Any) -> str: + return re.sub(r"\s+", " ", str(value).strip()).casefold() + + +def _extract_answer(response: str) -> tuple[str, bool]: + tagged = re.findall(r"\s*(.*?)\s*", response, flags=re.IGNORECASE | re.DOTALL) + if not tagged: + candidate = response.strip() + match = re.search(r"\b([A-D])\b", candidate, flags=re.IGNORECASE) + return (match.group(1).upper() if match else candidate), False + candidate = tagged[-1].strip() + match = re.search(r"\b([A-D])\b", candidate, flags=re.IGNORECASE) + return (match.group(1).upper() if match else candidate), True + + +def _score_sciknoweval(sample: Any) -> dict[str, Any]: + metadata = _metadata(sample) + expected = metadata.get("answer_key", getattr(sample, "label", "")) + task_type = str(metadata.get("task_type", "")).casefold() + response = getattr(sample, "response", "") + predicted, has_format = _extract_answer(response) + incorrect_format = int(not has_format) + + if "true_or_false" in task_type or "true/false" in task_type: + correct = _normalize(predicted) in {_normalize(expected), _normalize(str(expected).replace(" ", ""))} + else: + correct = _normalize(predicted) == _normalize(expected) + + if _was_truncated(sample): + feedback = "Your response was truncated because it exceeded the maximum length." + elif incorrect_format: + feedback = "Your answer had the wrong format. The solution must be given in the format: X." + else: + feedback = "" + return { + "score": 1.0 if correct and not incorrect_format else 0.0, + "predicted": predicted, + "format_error": incorrect_format, + "feedback": feedback, + } + + +def _extract_tool_calls(response: str) -> tuple[list[tuple[str, dict[str, Any]]], bool]: + """Parse (Action, Action Input) pairs in document order. + + Each ``Action Input:`` JSON block is paired with the most recent unmatched + ``Action:`` line; an Action with no following input parses as ``{}``. + """ + format_ok = bool(re.search(r"Action:.*?\nAction Input:", response, flags=re.IGNORECASE | re.DOTALL)) + events: list[tuple[int, str, Any]] = [] + for match in re.finditer(r"^\s*Action:\s*(.+?)\s*$", response, flags=re.IGNORECASE | re.MULTILINE): + events.append((match.start(), "action", match.group(1).strip())) + decoder = json.JSONDecoder() + for match in re.finditer(r"Action Input:\s*", response, flags=re.IGNORECASE): + try: + value, _ = decoder.raw_decode(response, match.end()) + except json.JSONDecodeError: + value = None + events.append((match.start(), "input", value if isinstance(value, dict) else None)) + events.sort(key=lambda event: event[0]) + pairs: list[tuple[str, dict[str, Any]]] = [] + pending_action: str | None = None + for _, kind, value in events: + if kind == "action": + if pending_action is not None: + pairs.append((pending_action, {})) + pending_action = value + else: + pairs.append((pending_action or "", value if value is not None else {})) + pending_action = None + if pending_action is not None: + pairs.append((pending_action, {})) + return pairs, format_ok + + +def _golden_tool_call(metadata: dict[str, Any]) -> list[tuple[str, dict[str, Any]]]: + golden = metadata.get("golden_answer") or [] + if isinstance(golden, str): + try: + golden = json.loads(golden) + except json.JSONDecodeError: + golden = [] + pairs: list[tuple[str, dict[str, Any]]] = [] + for item in golden: + if not isinstance(item, dict): + continue + action = str(item.get("Action", "")).strip() + value = item.get("Action_Input", {}) + if isinstance(value, str): + try: + value = json.loads(value) + except json.JSONDecodeError: + value = {} + if not isinstance(value, dict): + value = {} + pairs.append((action, value)) + return pairs + + +def _score_toolalpaca(sample: Any) -> dict[str, Any]: + metadata = _metadata(sample) + predicted_pairs, format_ok = _extract_tool_calls(getattr(sample, "response", "")) + expected_pairs = _golden_tool_call(metadata) + correct = format_ok and predicted_pairs == expected_pairs + if _was_truncated(sample): + feedback = "Your response was truncated because it exceeded the maximum length." + elif not format_ok: + feedback = "Use the required Action and Action Input format." + else: + feedback = "" + return { + "score": 1.0 if correct else 0.0, + "predicted_tool_calls": predicted_pairs, + "format_error": int(not format_ok), + "feedback": feedback, + } + + +def _score_one(sample: Any) -> dict[str, Any]: + source = str(_metadata(sample).get("data_source", "")).casefold() + if source == "sciknoweval": + return _score_sciknoweval(sample) + if source in {"toolalpaca", "tooluse"}: + return _score_toolalpaca(sample) + raise ValueError(f"Unsupported SDPO data_source: {source!r}") + + +def score(_args: Any, samples: Any) -> dict[str, Any] | list[dict[str, Any]]: + """Custom reward entry point for ``--custom-rm-path``. + + Relax calls this function with a list in ``--group-rm`` mode and with one + ``Sample`` otherwise. Returning only reward payloads keeps the reward + worker process isolated; the core rollout path later uses those payloads to + build dynamic teacher prompts. + """ + + if isinstance(samples, list): + return [_score_one(sample) for sample in samples] + return _score_one(samples) diff --git a/examples/on_policy_distillation/sdpo/run-sciknoweval-biology-4xgpu-colocate.sh b/examples/on_policy_distillation/sdpo/run-sciknoweval-biology-4xgpu-colocate.sh new file mode 100644 index 000000000..c5b42f803 --- /dev/null +++ b/examples/on_policy_distillation/sdpo/run-sciknoweval-biology-4xgpu-colocate.sh @@ -0,0 +1,132 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/../../.." +source scripts/models/qwen3-8B.sh +source examples/on_policy_distillation/sdpo/env.sh + +export CUDA_DEVICE_MAX_CONNECTIONS="${CUDA_DEVICE_MAX_CONNECTIONS:-1}" + +export RELAX_OPD_PER_POS_TOKEN_IDS=1 + +student_model="${STUDENT_MODEL_PATH:?Set STUDENT_MODEL_PATH}" +teacher_model="${TEACHER_MODEL_PATH:-${student_model}}" +data_path="${DATA_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/biology/train.jsonl}" +eval_path="${EVAL_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/biology/eval.jsonl}" +now="$(date '+%Y-%m-%d-%H:%M:%S')" +experiment_name="${EXPERIMENT_NAME:-sdpo-sciknoweval-biology-${now}}" + +CKPT_ARGS=( + --hf-checkpoint "${student_model}" + --megatron-to-hf-mode bridge + --attention-backend flash +) + +ROLLOUT_ARGS=( + --prompt-data "${data_path}" + --input-key prompt + --label-key label + --metadata-key metadata + --apply-chat-template + --group-rm + --custom-rm-path examples.on_policy_distillation.sdpo.reward.score + --reward-key score + --num-rollout 5000 + --rollout-batch-size 32 + --n-samples-per-prompt 8 + --global-batch-size 256 + --rollout-max-prompt-len 2048 + --rollout-max-response-len 8192 + --rollout-temperature 1.0 + --use-fault-tolerance +) + +EVAL_ARGS=( + --eval-interval 5 + --eval-prompt-data sciknoweval-biology "${eval_path}" + --n-samples-per-eval-prompt 16 +) + +OPD_ARGS=( + --use-opd + --opd-feedback-class "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" + --opd-type sglang + --teacher-hf-checkpoint "${teacher_model}" + --teacher-num-gpus-per-engine 1 + --teacher-sglang-mem-fraction-static 0.5 + --teacher-sglang-enable-weights-cpu-backup + --opd-loss-coef 1.0 + --opd-kl-coef 0.0 + --opd-disable-rl-reward + --opd-token-selection student_topk + --opd-log-prob-top-k 16 + --opd-kl-type jsd + --opd-jsd-alpha 0.5 + --opd-norm-mode tail + --opd-teacher-timeout-s 600 + --use-rollout-logprobs +) + +GRPO_ARGS=( + --advantage-estimator grpo + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.01 + --clip-grad 1.0 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --context-parallel-size 1 + --pipeline-model-parallel-size 1 + --calculate-per-token-loss + --no-masked-softmax-fusion + --optimizer-cpu-offload + --selective-offload + --use-precision-aware-optimizer + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + --qkv-format bshd + --micro-batch-size 1 +) + +SGLANG_ARGS=( + --rollout-num-gpus 3 + --rollout-num-gpus-per-engine 1 + --sglang-load-format dummy + --sglang-mem-fraction-static 0.45 +) + +MISC_ARGS=( + --resource '{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]}' + --max-staleness 0 + --num-data-storage-units 1 + --colocate + --train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}' + --use-health-check + --actor-num-gpus-per-node 4 + --num-gpus-per-node 4 + --tb-experiment-name "${experiment_name}" +) + +WANDB_ARGS=( + --use-wandb + --wandb-project relax-sdpo + --wandb-group "${WANDB_RUN_GROUP:-sdpo-sciknoweval-biology}" + --wandb-key "${WANDB_API_KEY}" +) + +exec python -m relax.entrypoints.train \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" "${ROLLOUT_ARGS[@]}" "${EVAL_ARGS[@]}" \ + "${OPD_ARGS[@]}" "${GRPO_ARGS[@]}" "${OPTIMIZER_ARGS[@]}" \ + "${PERF_ARGS[@]}" "${SGLANG_ARGS[@]}" "${MISC_ARGS[@]}" \ + "${WANDB_ARGS[@]}" diff --git a/examples/on_policy_distillation/sdpo/run-sciknoweval-chemistry-4xgpu-colocate.sh b/examples/on_policy_distillation/sdpo/run-sciknoweval-chemistry-4xgpu-colocate.sh new file mode 100644 index 000000000..292e212a4 --- /dev/null +++ b/examples/on_policy_distillation/sdpo/run-sciknoweval-chemistry-4xgpu-colocate.sh @@ -0,0 +1,132 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/../../.." +source scripts/models/qwen3-8B.sh +source examples/on_policy_distillation/sdpo/env.sh + +export CUDA_DEVICE_MAX_CONNECTIONS="${CUDA_DEVICE_MAX_CONNECTIONS:-1}" + +export RELAX_OPD_PER_POS_TOKEN_IDS=1 + +student_model="${STUDENT_MODEL_PATH:?Set STUDENT_MODEL_PATH}" +teacher_model="${TEACHER_MODEL_PATH:-${student_model}}" +data_path="${DATA_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/chemistry/train.jsonl}" +eval_path="${EVAL_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/chemistry/eval.jsonl}" +now="$(date '+%Y-%m-%d-%H:%M:%S')" +experiment_name="${EXPERIMENT_NAME:-sdpo-sciknoweval-chemistry-${now}}" + +CKPT_ARGS=( + --hf-checkpoint "${student_model}" + --megatron-to-hf-mode bridge + --attention-backend flash +) + +ROLLOUT_ARGS=( + --prompt-data "${data_path}" + --input-key prompt + --label-key label + --metadata-key metadata + --apply-chat-template + --group-rm + --custom-rm-path examples.on_policy_distillation.sdpo.reward.score + --reward-key score + --num-rollout 5000 + --rollout-batch-size 32 + --n-samples-per-prompt 8 + --global-batch-size 256 + --rollout-max-prompt-len 2048 + --rollout-max-response-len 8192 + --rollout-temperature 1.0 + --use-fault-tolerance +) + +EVAL_ARGS=( + --eval-interval 5 + --eval-prompt-data sciknoweval-chemistry "${eval_path}" + --n-samples-per-eval-prompt 16 +) + +OPD_ARGS=( + --use-opd + --opd-feedback-class "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" + --opd-type sglang + --teacher-hf-checkpoint "${teacher_model}" + --teacher-num-gpus-per-engine 1 + --teacher-sglang-mem-fraction-static 0.5 + --teacher-sglang-enable-weights-cpu-backup + --opd-loss-coef 1.0 + --opd-kl-coef 0.0 + --opd-disable-rl-reward + --opd-token-selection student_topk + --opd-log-prob-top-k 16 + --opd-kl-type jsd + --opd-jsd-alpha 0.5 + --opd-norm-mode tail + --opd-teacher-timeout-s 600 + --use-rollout-logprobs +) + +GRPO_ARGS=( + --advantage-estimator grpo + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.01 + --clip-grad 1.0 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --context-parallel-size 1 + --pipeline-model-parallel-size 1 + --calculate-per-token-loss + --no-masked-softmax-fusion + --optimizer-cpu-offload + --selective-offload + --use-precision-aware-optimizer + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + --qkv-format bshd + --micro-batch-size 1 +) + +SGLANG_ARGS=( + --rollout-num-gpus 3 + --rollout-num-gpus-per-engine 1 + --sglang-load-format dummy + --sglang-mem-fraction-static 0.45 +) + +MISC_ARGS=( + --resource '{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]}' + --max-staleness 0 + --num-data-storage-units 1 + --colocate + --train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}' + --use-health-check + --actor-num-gpus-per-node 4 + --num-gpus-per-node 4 + --tb-experiment-name "${experiment_name}" +) + +WANDB_ARGS=( + --use-wandb + --wandb-project relax-sdpo + --wandb-group "${WANDB_RUN_GROUP:-sdpo-sciknoweval-chemistry}" + --wandb-key "${WANDB_API_KEY}" +) + +exec python -m relax.entrypoints.train \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" "${ROLLOUT_ARGS[@]}" "${EVAL_ARGS[@]}" \ + "${OPD_ARGS[@]}" "${GRPO_ARGS[@]}" "${OPTIMIZER_ARGS[@]}" \ + "${PERF_ARGS[@]}" "${SGLANG_ARGS[@]}" "${MISC_ARGS[@]}" \ + "${WANDB_ARGS[@]}" diff --git a/examples/on_policy_distillation/sdpo/run-sciknoweval-material-4xgpu-colocate.sh b/examples/on_policy_distillation/sdpo/run-sciknoweval-material-4xgpu-colocate.sh new file mode 100644 index 000000000..5c63f2954 --- /dev/null +++ b/examples/on_policy_distillation/sdpo/run-sciknoweval-material-4xgpu-colocate.sh @@ -0,0 +1,132 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/../../.." +source scripts/models/qwen3-8B.sh +source examples/on_policy_distillation/sdpo/env.sh + +export CUDA_DEVICE_MAX_CONNECTIONS="${CUDA_DEVICE_MAX_CONNECTIONS:-1}" + +export RELAX_OPD_PER_POS_TOKEN_IDS=1 + +student_model="${STUDENT_MODEL_PATH:?Set STUDENT_MODEL_PATH}" +teacher_model="${TEACHER_MODEL_PATH:-${student_model}}" +data_path="${DATA_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/material/train.jsonl}" +eval_path="${EVAL_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/material/eval.jsonl}" +now="$(date '+%Y-%m-%d-%H:%M:%S')" +experiment_name="${EXPERIMENT_NAME:-sdpo-sciknoweval-material-${now}}" + +CKPT_ARGS=( + --hf-checkpoint "${student_model}" + --megatron-to-hf-mode bridge + --attention-backend flash +) + +ROLLOUT_ARGS=( + --prompt-data "${data_path}" + --input-key prompt + --label-key label + --metadata-key metadata + --apply-chat-template + --group-rm + --custom-rm-path examples.on_policy_distillation.sdpo.reward.score + --reward-key score + --num-rollout 5000 + --rollout-batch-size 32 + --n-samples-per-prompt 8 + --global-batch-size 256 + --rollout-max-prompt-len 2048 + --rollout-max-response-len 8192 + --rollout-temperature 1.0 + --use-fault-tolerance +) + +EVAL_ARGS=( + --eval-interval 5 + --eval-prompt-data sciknoweval-material "${eval_path}" + --n-samples-per-eval-prompt 16 +) + +OPD_ARGS=( + --use-opd + --opd-feedback-class "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" + --opd-type sglang + --teacher-hf-checkpoint "${teacher_model}" + --teacher-num-gpus-per-engine 1 + --teacher-sglang-mem-fraction-static 0.5 + --teacher-sglang-enable-weights-cpu-backup + --opd-loss-coef 1.0 + --opd-kl-coef 0.0 + --opd-disable-rl-reward + --opd-token-selection student_topk + --opd-log-prob-top-k 16 + --opd-kl-type jsd + --opd-jsd-alpha 0.5 + --opd-norm-mode tail + --opd-teacher-timeout-s 600 + --use-rollout-logprobs +) + +GRPO_ARGS=( + --advantage-estimator grpo + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.01 + --clip-grad 1.0 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --context-parallel-size 1 + --pipeline-model-parallel-size 1 + --calculate-per-token-loss + --no-masked-softmax-fusion + --optimizer-cpu-offload + --selective-offload + --use-precision-aware-optimizer + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + --qkv-format bshd + --micro-batch-size 1 +) + +SGLANG_ARGS=( + --rollout-num-gpus 3 + --rollout-num-gpus-per-engine 1 + --sglang-load-format dummy + --sglang-mem-fraction-static 0.45 +) + +MISC_ARGS=( + --resource '{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]}' + --max-staleness 0 + --num-data-storage-units 1 + --colocate + --train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}' + --use-health-check + --actor-num-gpus-per-node 4 + --num-gpus-per-node 4 + --tb-experiment-name "${experiment_name}" +) + +WANDB_ARGS=( + --use-wandb + --wandb-project relax-sdpo + --wandb-group "${WANDB_RUN_GROUP:-sdpo-sciknoweval-material}" + --wandb-key "${WANDB_API_KEY}" +) + +exec python -m relax.entrypoints.train \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" "${ROLLOUT_ARGS[@]}" "${EVAL_ARGS[@]}" \ + "${OPD_ARGS[@]}" "${GRPO_ARGS[@]}" "${OPTIMIZER_ARGS[@]}" \ + "${PERF_ARGS[@]}" "${SGLANG_ARGS[@]}" "${MISC_ARGS[@]}" \ + "${WANDB_ARGS[@]}" diff --git a/examples/on_policy_distillation/sdpo/run-sciknoweval-physics-4xgpu-colocate.sh b/examples/on_policy_distillation/sdpo/run-sciknoweval-physics-4xgpu-colocate.sh new file mode 100644 index 000000000..1b58dd03a --- /dev/null +++ b/examples/on_policy_distillation/sdpo/run-sciknoweval-physics-4xgpu-colocate.sh @@ -0,0 +1,132 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/../../.." +source scripts/models/qwen3-8B.sh +source examples/on_policy_distillation/sdpo/env.sh + +export CUDA_DEVICE_MAX_CONNECTIONS="${CUDA_DEVICE_MAX_CONNECTIONS:-1}" + +export RELAX_OPD_PER_POS_TOKEN_IDS=1 + +student_model="${STUDENT_MODEL_PATH:?Set STUDENT_MODEL_PATH}" +teacher_model="${TEACHER_MODEL_PATH:-${student_model}}" +data_path="${DATA_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/physics/train.jsonl}" +eval_path="${EVAL_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/sciknoweval/physics/eval.jsonl}" +now="$(date '+%Y-%m-%d-%H:%M:%S')" +experiment_name="${EXPERIMENT_NAME:-sdpo-sciknoweval-physics-${now}}" + +CKPT_ARGS=( + --hf-checkpoint "${student_model}" + --megatron-to-hf-mode bridge + --attention-backend flash +) + +ROLLOUT_ARGS=( + --prompt-data "${data_path}" + --input-key prompt + --label-key label + --metadata-key metadata + --apply-chat-template + --group-rm + --custom-rm-path examples.on_policy_distillation.sdpo.reward.score + --reward-key score + --num-rollout 5000 + --rollout-batch-size 32 + --n-samples-per-prompt 8 + --global-batch-size 256 + --rollout-max-prompt-len 2048 + --rollout-max-response-len 8192 + --rollout-temperature 1.0 + --use-fault-tolerance +) + +EVAL_ARGS=( + --eval-interval 5 + --eval-prompt-data sciknoweval-physics "${eval_path}" + --n-samples-per-eval-prompt 16 +) + +OPD_ARGS=( + --use-opd + --opd-feedback-class "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" + --opd-type sglang + --teacher-hf-checkpoint "${teacher_model}" + --teacher-num-gpus-per-engine 1 + --teacher-sglang-mem-fraction-static 0.5 + --teacher-sglang-enable-weights-cpu-backup + --opd-loss-coef 1.0 + --opd-kl-coef 0.0 + --opd-disable-rl-reward + --opd-token-selection student_topk + --opd-log-prob-top-k 16 + --opd-kl-type jsd + --opd-jsd-alpha 0.5 + --opd-norm-mode tail + --opd-teacher-timeout-s 600 + --use-rollout-logprobs +) + +GRPO_ARGS=( + --advantage-estimator grpo + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.01 + --clip-grad 1.0 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --context-parallel-size 1 + --pipeline-model-parallel-size 1 + --calculate-per-token-loss + --no-masked-softmax-fusion + --optimizer-cpu-offload + --selective-offload + --use-precision-aware-optimizer + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + --qkv-format bshd + --micro-batch-size 1 +) + +SGLANG_ARGS=( + --rollout-num-gpus 3 + --rollout-num-gpus-per-engine 1 + --sglang-load-format dummy + --sglang-mem-fraction-static 0.45 +) + +MISC_ARGS=( + --resource '{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]}' + --max-staleness 0 + --num-data-storage-units 1 + --colocate + --train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}' + --use-health-check + --actor-num-gpus-per-node 4 + --num-gpus-per-node 4 + --tb-experiment-name "${experiment_name}" +) + +WANDB_ARGS=( + --use-wandb + --wandb-project relax-sdpo + --wandb-group "${WANDB_RUN_GROUP:-sdpo-sciknoweval-physics}" + --wandb-key "${WANDB_API_KEY}" +) + +exec python -m relax.entrypoints.train \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" "${ROLLOUT_ARGS[@]}" "${EVAL_ARGS[@]}" \ + "${OPD_ARGS[@]}" "${GRPO_ARGS[@]}" "${OPTIMIZER_ARGS[@]}" \ + "${PERF_ARGS[@]}" "${SGLANG_ARGS[@]}" "${MISC_ARGS[@]}" \ + "${WANDB_ARGS[@]}" diff --git a/examples/on_policy_distillation/sdpo/run-toolalpaca-4xgpu-colocate.sh b/examples/on_policy_distillation/sdpo/run-toolalpaca-4xgpu-colocate.sh new file mode 100644 index 000000000..1a055b8a5 --- /dev/null +++ b/examples/on_policy_distillation/sdpo/run-toolalpaca-4xgpu-colocate.sh @@ -0,0 +1,132 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/../../.." +source scripts/models/qwen3-8B.sh +source examples/on_policy_distillation/sdpo/env.sh + +export CUDA_DEVICE_MAX_CONNECTIONS="${CUDA_DEVICE_MAX_CONNECTIONS:-1}" + +export RELAX_OPD_PER_POS_TOKEN_IDS=1 + +student_model="${STUDENT_MODEL_PATH:?Set STUDENT_MODEL_PATH}" +teacher_model="${TEACHER_MODEL_PATH:-${student_model}}" +data_path="${DATA_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/toolalpaca/train.jsonl}" +eval_path="${EVAL_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/toolalpaca/eval.jsonl}" +now="$(date '+%Y-%m-%d-%H:%M:%S')" +experiment_name="${EXPERIMENT_NAME:-sdpo-toolalpaca-${now}}" + +CKPT_ARGS=( + --hf-checkpoint "${student_model}" + --megatron-to-hf-mode bridge + --attention-backend flash +) + +ROLLOUT_ARGS=( + --prompt-data "${data_path}" + --input-key prompt + --label-key label + --metadata-key metadata + --apply-chat-template + --group-rm + --custom-rm-path examples.on_policy_distillation.sdpo.reward.score + --reward-key score + --num-rollout 5000 + --rollout-batch-size 32 + --n-samples-per-prompt 8 + --global-batch-size 256 + --rollout-max-prompt-len 2048 + --rollout-max-response-len 8192 + --rollout-temperature 1.0 + --use-fault-tolerance +) + +EVAL_ARGS=( + --eval-interval 5 + --eval-prompt-data toolalpaca "${eval_path}" + --n-samples-per-eval-prompt 16 +) + +OPD_ARGS=( + --use-opd + --opd-feedback-class "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" + --opd-type sglang + --teacher-hf-checkpoint "${teacher_model}" + --teacher-num-gpus-per-engine 1 + --teacher-sglang-mem-fraction-static 0.5 + --teacher-sglang-enable-weights-cpu-backup + --opd-loss-coef 1.0 + --opd-kl-coef 0.0 + --opd-disable-rl-reward + --opd-token-selection student_topk + --opd-log-prob-top-k 16 + --opd-kl-type jsd + --opd-jsd-alpha 0.5 + --opd-norm-mode tail + --opd-teacher-timeout-s 600 + --use-rollout-logprobs +) + +GRPO_ARGS=( + --advantage-estimator grpo + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.01 + --clip-grad 1.0 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --context-parallel-size 1 + --pipeline-model-parallel-size 1 + --calculate-per-token-loss + --no-masked-softmax-fusion + --optimizer-cpu-offload + --selective-offload + --use-precision-aware-optimizer + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + --qkv-format bshd + --micro-batch-size 1 +) + +SGLANG_ARGS=( + --rollout-num-gpus 3 + --rollout-num-gpus-per-engine 1 + --sglang-load-format dummy + --sglang-mem-fraction-static 0.45 +) + +MISC_ARGS=( + --resource '{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]}' + --max-staleness 0 + --num-data-storage-units 1 + --colocate + --train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}' + --use-health-check + --actor-num-gpus-per-node 4 + --num-gpus-per-node 4 + --tb-experiment-name "${experiment_name}" +) + +WANDB_ARGS=( + --use-wandb + --wandb-project relax-sdpo + --wandb-group "${WANDB_RUN_GROUP:-sdpo-toolalpaca}" + --wandb-key "${WANDB_API_KEY}" +) + +exec python -m relax.entrypoints.train \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" "${ROLLOUT_ARGS[@]}" "${EVAL_ARGS[@]}" \ + "${OPD_ARGS[@]}" "${GRPO_ARGS[@]}" "${OPTIMIZER_ARGS[@]}" \ + "${PERF_ARGS[@]}" "${SGLANG_ARGS[@]}" "${MISC_ARGS[@]}" \ + "${WANDB_ARGS[@]}" diff --git a/examples/on_policy_distillation/sdpo/run-tooluse-4xgpu-colocate.sh b/examples/on_policy_distillation/sdpo/run-tooluse-4xgpu-colocate.sh new file mode 100644 index 000000000..b3db0145c --- /dev/null +++ b/examples/on_policy_distillation/sdpo/run-tooluse-4xgpu-colocate.sh @@ -0,0 +1,132 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +set -euo pipefail +cd "$(dirname "${BASH_SOURCE[0]}")/../../.." +source scripts/models/qwen3-8B.sh +source examples/on_policy_distillation/sdpo/env.sh + +export CUDA_DEVICE_MAX_CONNECTIONS="${CUDA_DEVICE_MAX_CONNECTIONS:-1}" + +export RELAX_OPD_PER_POS_TOKEN_IDS=1 + +student_model="${STUDENT_MODEL_PATH:?Set STUDENT_MODEL_PATH}" +teacher_model="${TEACHER_MODEL_PATH:-${student_model}}" +data_path="${DATA_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/tooluse/train.jsonl}" +eval_path="${EVAL_PATH:-${SDPO_DATA_ROOT:?Set SDPO_DATA_ROOT}/tooluse/eval.jsonl}" +now="$(date '+%Y-%m-%d-%H:%M:%S')" +experiment_name="${EXPERIMENT_NAME:-sdpo-tooluse-${now}}" + +CKPT_ARGS=( + --hf-checkpoint "${student_model}" + --megatron-to-hf-mode bridge + --attention-backend flash +) + +ROLLOUT_ARGS=( + --prompt-data "${data_path}" + --input-key prompt + --label-key label + --metadata-key metadata + --apply-chat-template + --group-rm + --custom-rm-path examples.on_policy_distillation.sdpo.reward.score + --reward-key score + --num-rollout 5000 + --rollout-batch-size 32 + --n-samples-per-prompt 8 + --global-batch-size 256 + --rollout-max-prompt-len 2048 + --rollout-max-response-len 8192 + --rollout-temperature 1.0 + --use-fault-tolerance +) + +EVAL_ARGS=( + --eval-interval 5 + --eval-prompt-data tooluse "${eval_path}" + --n-samples-per-eval-prompt 16 +) + +OPD_ARGS=( + --use-opd + --opd-feedback-class "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" + --opd-type sglang + --teacher-hf-checkpoint "${teacher_model}" + --teacher-num-gpus-per-engine 1 + --teacher-sglang-mem-fraction-static 0.5 + --teacher-sglang-enable-weights-cpu-backup + --opd-loss-coef 1.0 + --opd-kl-coef 0.0 + --opd-disable-rl-reward + --opd-token-selection student_topk + --opd-log-prob-top-k 16 + --opd-kl-type jsd + --opd-jsd-alpha 0.5 + --opd-norm-mode tail + --opd-teacher-timeout-s 600 + --use-rollout-logprobs +) + +GRPO_ARGS=( + --advantage-estimator grpo + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.01 + --clip-grad 1.0 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --context-parallel-size 1 + --pipeline-model-parallel-size 1 + --calculate-per-token-loss + --no-masked-softmax-fusion + --optimizer-cpu-offload + --selective-offload + --use-precision-aware-optimizer + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + --qkv-format bshd + --micro-batch-size 1 +) + +SGLANG_ARGS=( + --rollout-num-gpus 3 + --rollout-num-gpus-per-engine 1 + --sglang-load-format dummy + --sglang-mem-fraction-static 0.45 +) + +MISC_ARGS=( + --resource '{"actor": [1, 4], "rollout": [1, 3], "teacher": [1, 1]}' + --max-staleness 0 + --num-data-storage-units 1 + --colocate + --train-env-vars '{"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True"}' + --use-health-check + --actor-num-gpus-per-node 4 + --num-gpus-per-node 4 + --tb-experiment-name "${experiment_name}" +) + +WANDB_ARGS=( + --use-wandb + --wandb-project relax-sdpo + --wandb-group "${WANDB_RUN_GROUP:-sdpo-tooluse}" + --wandb-key "${WANDB_API_KEY}" +) + +exec python -m relax.entrypoints.train \ + "${MODEL_ARGS[@]}" \ + "${CKPT_ARGS[@]}" "${ROLLOUT_ARGS[@]}" "${EVAL_ARGS[@]}" \ + "${OPD_ARGS[@]}" "${GRPO_ARGS[@]}" "${OPTIMIZER_ARGS[@]}" \ + "${PERF_ARGS[@]}" "${SGLANG_ARGS[@]}" "${MISC_ARGS[@]}" \ + "${WANDB_ARGS[@]}" diff --git a/relax/backends/megatron/actor.py b/relax/backends/megatron/actor.py index 6d9eae24a..f8666cbee 100644 --- a/relax/backends/megatron/actor.py +++ b/relax/backends/megatron/actor.py @@ -1298,8 +1298,8 @@ def train_hybrid(self, rollout_id) -> None: data_fields += ["rollout_routed_experts"] if self.args.use_rollout_routing_replay else [] if self.args.multimodal_keys is not None: data_fields.append("multimodal_train_inputs") - if self.args.use_opd and self.args.opd_type == "sglang": - data_fields.append("teacher_log_probs") + if self.args.use_opd: + consume_opd_train_data(data_fields, self.args) with timer("train_get_data"): sub_batch, batch_meta = self._get_data_from_transfer_queue( "train", rollout_id, data_fields, batch_size, batch_index diff --git a/relax/backends/megatron/data.py b/relax/backends/megatron/data.py index 4fd2a6c89..1dec33e6c 100644 --- a/relax/backends/megatron/data.py +++ b/relax/backends/megatron/data.py @@ -21,7 +21,7 @@ from relax.utils.data.seqlen_balancing import get_seqlen_balanced_partitions from relax.utils.logging_utils import get_logger from relax.utils.metrics.metric_utils import compute_rollout_step -from relax.utils.opd.opd_utils import OPD_ROLLOUT_LOG_SKIP_FIELDS +from relax.utils.opd.opd_utils import OPD_ROLLOUT_LOG_SKIP_FIELDS, OPD_SAMPLE_MASK from relax.utils.timer import Timer from relax.utils.training import train_metric_utils from relax.utils.training.flops_counter import FlopsCounter @@ -292,6 +292,9 @@ def get_batch( else: batch, _ = next(data_iterator) + if getattr(get_args(), "use_opd", False): + _apply_opd_sample_mask(batch) + use_dynamic_context_parallel = getattr(get_args(), "dynamic_context_parallel", False) if use_dynamic_context_parallel: # Pick this mb's CP size with the SAME per-GPU token budget the iterator was packed @@ -586,6 +589,28 @@ def get_batch( return batch +def _apply_opd_sample_mask(batch: dict) -> None: + """Fold the OPD sample mask into the per-sample response loss masks. + + Inactive samples get an all-zero loss mask, so the standard loss reduction + (num_tokens / sum_of_sample_mean / CP slicing / metric counts) excludes + them from both the numerator and the denominator without any extra + plumbing. + """ + sample_mask = batch.get(OPD_SAMPLE_MASK) + loss_masks = batch.get("loss_masks") + if sample_mask is None or loss_masks is None: + return + if len(sample_mask) != len(loss_masks): + raise ValueError( + f"{OPD_SAMPLE_MASK} must contain one scalar per sample: expected {len(loss_masks)}, got {len(sample_mask)}" + ) + batch["loss_masks"] = [ + torch.as_tensor(loss_mask, dtype=torch.float32) * float(mask) + for loss_mask, mask in zip(loss_masks, sample_mask, strict=True) + ] + + def gather_log_data( metric_name: str, args: Namespace, diff --git a/relax/backends/megatron/model.py b/relax/backends/megatron/model.py index 477ba9194..cf6e5a7f0 100644 --- a/relax/backends/megatron/model.py +++ b/relax/backends/megatron/model.py @@ -1260,7 +1260,7 @@ def forward_step( # already counted once. A `* cp_size` here would over-weight metrics by # CP degree under dynamic CP (and is a no-op under static CP, where the # count previously carried the cancelling cp factor). - loss_reduced[key] = value / num_samples_or_tokens + loss_reduced[key] = value / num_samples_or_tokens if num_samples_or_tokens != 0 else 0.0 capture_hooks.end_step_for() return loss_reduced, grad_norm capture_hooks.end_step_for() diff --git a/relax/engine/rollout/data_source.py b/relax/engine/rollout/data_source.py index ad790216e..4f78a3396 100644 --- a/relax/engine/rollout/data_source.py +++ b/relax/engine/rollout/data_source.py @@ -67,6 +67,8 @@ def _create_dataset(args, tokenizer, processor, multimodal_config=None): teacher_mm_keys = {"image": _t_img} logger.info(f"OPSD: teacher_multimodal_keys = {teacher_mm_keys}") + teacher_prompt_key = (getattr(args, "opd_feedback_kwargs", None) or {}).get("teacher_prompt_key") + if use_streaming: from relax.utils.data.streaming_dataset import StreamingDataset @@ -101,7 +103,7 @@ def _create_dataset(args, tokenizer, processor, multimodal_config=None): prefetch_num_workers=prefetch_num_workers, multimodal_config=multimodal_config, custom_prompt_func=custom_prompt_func, - teacher_prompt_key=getattr(args, "opd_teacher_prompt_key", None), + teacher_prompt_key=teacher_prompt_key, teacher_multimodal_keys=teacher_mm_keys, ) else: @@ -123,7 +125,7 @@ def _create_dataset(args, tokenizer, processor, multimodal_config=None): seed=args.rollout_seed, multimodal_config=multimodal_config, custom_prompt_func=custom_prompt_func, - teacher_prompt_key=getattr(args, "opd_teacher_prompt_key", None), + teacher_prompt_key=teacher_prompt_key, teacher_multimodal_keys=teacher_mm_keys, ) diff --git a/relax/engine/rollout/on_policy_distillation.py b/relax/engine/rollout/on_policy_distillation.py index b888b4210..c24af7df3 100644 --- a/relax/engine/rollout/on_policy_distillation.py +++ b/relax/engine/rollout/on_policy_distillation.py @@ -3,12 +3,14 @@ import asyncio import json from collections.abc import Awaitable, Callable, Sequence +from time import monotonic import aiohttp import numpy as np from relax.utils.logging_utils import get_logger from relax.utils.opd import opd_main_worker, opd_opsd_worker +from relax.utils.opd.feedback import load_feedback from relax.utils.types import Sample @@ -95,9 +97,15 @@ def _pick_teacher_url(args, sample=None) -> str: class OpdManager: def __init__(self, args): self.args = args + self.feedback = load_feedback( + getattr(args, "opd_feedback_class", None), getattr(args, "opd_feedback_kwargs", None) + ) self.topk_worker: opd_main_worker.TopkWorker | None = None self.sampled_worker: opd_main_worker.SampledTokenWorker | None = None # 仅 student_sampled self.opsd_worker: opd_opsd_worker.OpsdWorker | None = None + opsd_worker = self.feedback.create_opsd_worker(args) + if opsd_worker.is_opsd: + self.opsd_worker = opsd_worker token_selection = args.opd_token_selection if token_selection != "student_sampled": @@ -105,9 +113,6 @@ def __init__(self, args): else: self.sampled_worker = opd_main_worker.SampledTokenWorker.from_args(args) - opsd_worker = opd_opsd_worker.OpsdWorker.from_args(args) - self.opsd_worker = opsd_worker if opsd_worker.is_opsd else None - @property def is_topk(self) -> bool: return self.topk_worker is not None @@ -117,7 +122,7 @@ def is_opsd(self) -> bool: return self.opsd_worker is not None def schema_opd_transfer_data(self) -> list[str]: - fields: list[str] = [] + fields: list[str] = list(self.feedback.extra_transfer_schema()) if self.topk_worker is not None: fields.extend(self.topk_worker.topk_transfer_fields()) if self.sampled_worker is not None: @@ -125,8 +130,9 @@ def schema_opd_transfer_data(self) -> list[str]: return fields def produce_opd_transfer_data(self, samples: list[Sample], train_data: dict) -> None: + self.feedback.produce_extra_transfer(samples, train_data) if self.topk_worker is not None: - for field_name in opd_main_worker.TopkWorker.TRANSFER_FIELDS: + for field_name in self.topk_worker.topk_transfer_fields(): if not any(getattr(s, field_name, None) is not None for s in samples): continue flat: list = [] @@ -137,11 +143,6 @@ def produce_opd_transfer_data(self, samples: list[Sample], train_data: dict) -> else: flat.append(v.reshape(-1).tolist()) train_data[field_name] = flat - kl_field = opd_main_worker.TopkWorker.TRANSFER_K_LENGTHS - if self.topk_worker.spec.name == "union" and any(getattr(s, kl_field, None) is not None for s in samples): - train_data[kl_field] = [ - getattr(s, kl_field).tolist() if getattr(s, kl_field, None) is not None else [] for s in samples - ] elif self.sampled_worker is not None: train_data[opd_main_worker.SampledTokenWorker.TRANSFER_TEACHER_LOG_PROBS] = [ s.teacher_log_probs if s.teacher_log_probs is not None else [] for s in samples @@ -161,7 +162,11 @@ def parse_rollout_logprobs(self, meta_info: dict, tokens: list, log_probs: list) if val_b64 is None: return tokens, log_probs import numpy as np - import pybase64 + + try: + import pybase64 + except ImportError: # pragma: no cover - pybase64 is used in the runtime image + import base64 as pybase64 val = np.frombuffer(pybase64.b64decode(val_b64), dtype=np.float32) idx_b64 = meta_info.get("output_token_logprobs_idx_b64") @@ -197,18 +202,29 @@ async def prefill( sample_list = list(samples) if isinstance(samples, Sequence) else [samples] if self.opsd_worker is not None: - await asyncio.gather(*[self.opsd_worker.build_teacher_inputs(self.args, s) for s in sample_list]) - - async with _create_teacher_client_session(self.args) as session: - fetch_results = await asyncio.gather(*[self._teacher_prefill(s, session) for s in sample_list]) - self._raise_if_all_failed(sample_list, fetch_results) - - if self.topk_worker is not None and self.topk_worker.spec.student_at_teacher: - await asyncio.gather( - *[self._student_prefill(s, session, encode_multimodal_inputs) for s in sample_list] + await asyncio.gather(*[self.opsd_worker.build_teacher_inputs(self.args, sample) for sample in sample_list]) + + prefill_start = monotonic() + fetch_results: list[bool] = [] + if sample_list: + async with _create_teacher_client_session(self.args) as session: + fetch_results = list( + await asyncio.gather(*[self._teacher_prefill(sample, session) for sample in sample_list]) ) + self._raise_if_all_failed(sample_list, fetch_results) + + if self.topk_worker is not None and self.topk_worker.spec.student_at_teacher: + await asyncio.gather( + *[self._student_prefill(sample, session, encode_multimodal_inputs) for sample in sample_list] + ) self._assemble_transfer(sample_list) + logger.info( + "OPD teacher prefill: requests=%d valid=%d elapsed=%.3fs", + len(sample_list), + sum(int(result is True) for result in fetch_results), + monotonic() - prefill_start, + ) async def _post_logprob( self, @@ -237,25 +253,23 @@ async def _post_logprob( return opd_main_worker.LogprobResponse(data) async def _teacher_prefill(self, sample: Sample, session: aiohttp.ClientSession) -> bool: - from relax.utils.opd.opd_utils import build_teacher_preexpanded_image_data - response_length = int(sample.response_length or 0) if response_length <= 0: return True - # OPSD: expanded teacher_tokens + image_data; if self.opsd_worker is not None: - image_data = await build_teacher_preexpanded_image_data(sample) + image_data = self.opsd_worker.build_preexpanded_image_data(sample) teacher_input_ids = self.opsd_worker.teacher_input_ids(sample, response_length) prompt_length = self.opsd_worker.teacher_prompt_len(sample, response_length) + logprob_start_len = max(prompt_length - 1, 0) else: image_data = None teacher_input_ids = sample.rollout_tokens or sample.tokens - prompt_length = len(sample.tokens) - response_length - logprob_start_len = max(prompt_length - 1, 0) + logprob_start_len = max(len(sample.tokens) - response_length - 1, 0) mm_fields = {"image_data": image_data} if image_data is not None else None if self.topk_worker is not None: + self.feedback.check_student_topk_ids(sample, self.topk_worker.top_k) payload = self.topk_worker.build_teacher_payload( input_ids=teacher_input_ids, logprob_start_len=logprob_start_len, @@ -361,6 +375,7 @@ def _assemble_transfer(self, samples: list[Sample]) -> None: teacher_at_student_lp=sample.teacher_at_student_topk_log_probs, student_at_teacher_lp=sample.student_at_teacher_topk_log_probs, ) + self.feedback.check_transfer_channels(sample, channels, self.topk_worker.top_k) sample.opd_topk_token_ids = channels.get(opd_main_worker.TopkWorker.TRANSFER_TOKEN_IDS) sample.opd_topk_student_log_probs = channels.get(opd_main_worker.TopkWorker.TRANSFER_STUDENT_LOG_PROBS) sample.opd_topk_teacher_log_probs = channels.get(opd_main_worker.TopkWorker.TRANSFER_TEACHER_LOG_PROBS) diff --git a/relax/engine/rollout/sglang_rollout.py b/relax/engine/rollout/sglang_rollout.py index 0a0149cda..bbaab2c4e 100644 --- a/relax/engine/rollout/sglang_rollout.py +++ b/relax/engine/rollout/sglang_rollout.py @@ -6,7 +6,7 @@ import inspect import uuid from argparse import Namespace -from collections.abc import Callable +from collections.abc import Awaitable, Callable from contextlib import AbstractAsyncContextManager, contextmanager from time import monotonic from typing import Any @@ -587,7 +587,7 @@ async def generate_and_rm( sample.reward = reward if state.opd_manager and not evaluation: - await state.opd_manager.prefill(samples, _encode_multimodal_inputs) + await _record_feedback_and_prefill(state, samples, _encode_multimodal_inputs) return samples else: @@ -598,7 +598,7 @@ async def generate_and_rm( sample.reward = await async_rm(args, sample) if state.opd_manager and not evaluation: - await state.opd_manager.prefill(sample, _encode_multimodal_inputs) + await _record_feedback_and_prefill(state, [sample], _encode_multimodal_inputs) return sample @@ -634,6 +634,21 @@ def _aggregate_rollout_timing(all_samples: list[Sample], get_samples_times: list return metrics +async def _record_feedback_and_prefill( + state: GenerateState, + group: list[Sample], + encode_multimodal_inputs: Callable[[dict], Awaitable[tuple[dict, float]]], +) -> None: + """The shared OPD-family two-step: record env feedback, then build the + privileged teacher prompts, before the teacher prefill fetch.""" + feedback = state.opd_manager.feedback + rewards = [sample.reward for sample in group] + for sample, reward in zip(group, rewards): + feedback.record_sample_feedback(sample, reward) + feedback.prepare_teacher_prompts(group, rewards) + await state.opd_manager.prefill(group, encode_multimodal_inputs) + + async def generate_and_rm_group( args: Namespace, group: list[Sample], sampling_params: dict[str, Any], evaluation: bool = False ) -> list[Sample]: @@ -678,7 +693,7 @@ async def generate_and_rm_group( sample.reward = reward if state.opd_manager and not evaluation: - await state.opd_manager.prefill(group, _encode_multimodal_inputs) + await _record_feedback_and_prefill(state, group, _encode_multimodal_inputs) return group diff --git a/relax/utils/data/data_utils.py b/relax/utils/data/data_utils.py index d1a4fd84f..f220eb369 100644 --- a/relax/utils/data/data_utils.py +++ b/relax/utils/data/data_utils.py @@ -427,12 +427,16 @@ def process_raw_sample( build_messages_fn=build_messages, ) + # The rendered teacher prompt is surfaced as raw metadata; OPSDFeedback owns + # the policy of assigning it to sample.teacher_prompt at rollout time. + if teacher_prompt_str is not None and isinstance(metadata, dict): + metadata["opd_teacher_prompt"] = teacher_prompt_str + return Sample( prompt=output_prompt, label=data[label_key] if label_key is not None else None, metadata=metadata, multimodal_inputs=multimodal_inputs, - teacher_prompt=teacher_prompt_str, teacher_multimodal_inputs=teacher_multimodal_inputs, ) diff --git a/relax/utils/opd/feedback.py b/relax/utils/opd/feedback.py new file mode 100644 index 000000000..31664979f --- /dev/null +++ b/relax/utils/opd/feedback.py @@ -0,0 +1,117 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""Feedback strategies that own every OPD-flavored algorithm difference. + +``EnvironmentFeedback`` is the polymorphism point of the OPD algorithm family +(OPD / MOPD / OPSD / SDPO). Every rollout runs the same shared two-step path +-- record what the environment said, then decide what privileged prompt the +teacher sees -- bound dynamically to a subclass. The base-class defaults +reproduce plain OPD (empty privilege), so a subclass only overrides the hooks +where its behavior differs. Hooks raise to hard-fail and return to degrade. + +The constructor parameter schema is declared once on the interface; ``args`` +carries it verbatim via ``--opd-feedback-kwargs`` and the selected subclass +binds it at construction time, consuming only the fields it uses. +""" + +from __future__ import annotations + +import importlib +from typing import Any + +from relax.utils.opd.opd_opsd_worker import OpsdWorker +from relax.utils.types import Sample + + +class EnvironmentFeedback: + def __init__(self, teacher_prompt_key: str | None = None, success_reward_threshold: float = 1.0) -> None: + self.teacher_prompt_key = teacher_prompt_key + self.success_reward_threshold = success_reward_threshold + + @staticmethod + def record(sample: Sample, text: str | None) -> None: + if text: + sample.metadata.setdefault("env_feedback", []).append(str(text)) + + def record_sample_feedback(self, sample: Sample, reward: Any) -> None: + feedback = self._reward_feedback(reward) + if feedback: + self.record(sample, feedback) + + def prepare_teacher_prompts(self, group: list[Sample], rewards: list[Any]) -> None: + """Assign the privileged teacher prompt for every sample in the + group.""" + for sample in group: + sample.teacher_prompt = None + sample.opd_sample_mask = None + + @staticmethod + def create_opsd_worker(args: Any) -> OpsdWorker: + return OpsdWorker.from_args(args) + + @classmethod + def validate_launch_args(cls, args: Any) -> None: + return None + + def extra_transfer_schema(self) -> list[str]: + return [] + + def produce_extra_transfer(self, samples: list[Sample], train_data: dict) -> None: + return None + + def check_student_topk_ids(self, sample: Sample, top_k: int) -> None: + return None + + def check_transfer_channels(self, sample: Sample, channels: dict, top_k: int) -> None: + return None + + @staticmethod + def _reward_feedback(reward: Any) -> str | None: + if isinstance(reward, dict): + for key in ("feedback", "error", "feedback_raw"): + value = reward.get(key) + if value: + if isinstance(value, str): + return value if value.strip() else None + return str(value) + return None + + @staticmethod + def feedback_text(sample: Sample, reward: Any) -> str: + reward_feedback = EnvironmentFeedback._reward_feedback(reward) + if reward_feedback: + return reward_feedback + values = sample.metadata.get("env_feedback", []) if isinstance(sample.metadata, dict) else [] + return "\n".join(str(value) for value in values if value) + + +class OPDFeedback(EnvironmentFeedback): + """Plain OPD/MOPD: the teacher sees exactly what the student saw.""" + + +def load_feedback_class(path: str | None) -> type[EnvironmentFeedback]: + if not path: + return OPDFeedback + module_name, separator, class_name = path.rpartition(".") + if not separator: + raise ValueError(f"Invalid feedback class path: {path!r}") + cls = getattr(importlib.import_module(module_name), class_name, None) + if not isinstance(cls, type) or not issubclass(cls, EnvironmentFeedback): + raise TypeError(f"{path!r} must name an EnvironmentFeedback subclass") + return cls + + +def load_feedback(path: str | None, kwargs: dict[str, Any] | None) -> EnvironmentFeedback: + cls = load_feedback_class(path) + try: + return cls(**(kwargs or {})) + except TypeError as e: + raise TypeError(f"Invalid --opd-feedback-kwargs for {cls.__name__}: {kwargs or {}} ({e})") from e + + +__all__ = [ + "EnvironmentFeedback", + "OPDFeedback", + "load_feedback", + "load_feedback_class", +] diff --git a/relax/utils/opd/opd_main_worker.py b/relax/utils/opd/opd_main_worker.py index 3d1843828..eb57dd75f 100644 --- a/relax/utils/opd/opd_main_worker.py +++ b/relax/utils/opd/opd_main_worker.py @@ -5,7 +5,12 @@ from dataclasses import dataclass, replace import numpy as np -import pybase64 + + +try: + import pybase64 +except ImportError: # pragma: no cover - pybase64 is used in the runtime image + import base64 as pybase64 @dataclass @@ -39,17 +44,42 @@ def _decode_topk_2d(self, prefix: str, response_length: int | None, top_k: int): if top_k <= 0: return None val = self._b64_decode(self.meta.get(f"{prefix}_val_b64"), "float32") - n = val.size // top_k - if n <= 0: + if val.size: + n = val.size // top_k + if n <= 0: + return None + if response_length is None: + response_length = n + if response_length <= 0 or n < response_length: + return None + take = response_length * top_k + lps = val[-take:].reshape(response_length, top_k) + idx = self._b64_decode(self.meta.get(f"{prefix}_idx_b64"), "int32") + ids = idx[-take:].reshape(response_length, top_k) if idx.size >= take else None + return ids, lps + + rows = self.meta.get(prefix) + if not isinstance(rows, list) or not rows: return None if response_length is None: - response_length = n - if response_length <= 0 or n < response_length: + response_length = len(rows) + if response_length <= 0 or len(rows) < response_length: + return None + + selected_rows = rows[-response_length:] + if any(row is None or len(row) < top_k for row in selected_rows): + return None + try: + lps = np.asarray( + [[float(item[0]) for item in row[:top_k]] for row in selected_rows], + dtype=np.float32, + ) + ids = np.asarray( + [[int(item[1]) for item in row[:top_k]] for row in selected_rows], + dtype=np.int32, + ) + except (IndexError, TypeError, ValueError): return None - take = response_length * top_k - lps = val[-take:].reshape(response_length, top_k) - idx = self._b64_decode(self.meta.get(f"{prefix}_idx_b64"), "int32") - ids = idx[-take:].reshape(response_length, top_k) if idx.size >= take else None return ids, lps def base_logprobs_1d(self) -> np.ndarray | None: @@ -101,7 +131,6 @@ class TopkWorker: TRANSFER_STUDENT_LOG_PROBS = "opd_topk_student_log_probs" # only union TRANSFER_K_LENGTHS = "opd_topk_ksz" - TRANSFER_FIELDS = (TRANSFER_TOKEN_IDS, TRANSFER_STUDENT_LOG_PROBS, TRANSFER_TEACHER_LOG_PROBS) def __init__( self, @@ -120,7 +149,7 @@ def __init__( self._teacher_prefill_tpl: dict = {} if self.spec.teacher_self_topk: self._teacher_prefill_tpl["top_logprobs_num"] = self.top_k - self._student_prefill_tpl: dict = {} if self.spec.student_at_teacher else {} + self._student_prefill_tpl: dict = {} @classmethod def from_args(cls, args) -> "TopkWorker": diff --git a/relax/utils/opd/opd_opsd_worker.py b/relax/utils/opd/opd_opsd_worker.py index af3a1d5c3..19acdc8e6 100644 --- a/relax/utils/opd/opd_opsd_worker.py +++ b/relax/utils/opd/opd_opsd_worker.py @@ -20,9 +20,8 @@ def __init__(self, *, is_opsd: bool = False): @classmethod def from_args(cls, args) -> "OpsdWorker": - teacher_prompt_key = getattr(args, "opd_teacher_prompt_key", None) teacher_image_key = getattr(args, "opd_teacher_image_key", None) - return cls(is_opsd=teacher_prompt_key is not None or teacher_image_key is not None) + return cls(is_opsd=teacher_image_key is not None) async def build_teacher_inputs(self, args, sample: "Sample") -> None: """Pre-expand teacher inputs on the client (rollout) side. diff --git a/relax/utils/opd/opd_utils.py b/relax/utils/opd/opd_utils.py index ce052263a..72deaced7 100644 --- a/relax/utils/opd/opd_utils.py +++ b/relax/utils/opd/opd_utils.py @@ -12,7 +12,6 @@ import torch.distributed as dist from relax.utils.logging_utils import get_logger -from relax.utils.opd import opd_opsd_worker if TYPE_CHECKING: @@ -21,6 +20,8 @@ logger = get_logger(__name__) +PREEXPANDED_RAW_FORMAT = "opd_preexpanded_raw" + OPD_TOKEN_SELECTIONS = ("student_sampled", "student_topk", "teacher_topk", "union") OPD_KL_TYPES = ("reverse_kl", "forward_kl", "low_var_kl", "jsd") @@ -34,6 +35,7 @@ ) OPD_CP_FLOAT_FIELDS = ("teacher_log_probs",) +OPD_SAMPLE_MASK = "opd_sample_mask" def iter_opd_cp_float_fields() -> tuple[str, ...]: @@ -596,7 +598,11 @@ def add_opd_arguments(parser: Any) -> Any: "--opd-jsd-alpha", type=float, default=0.5, - help="Mixture coefficient for --opd-kl-type=jsd. 0.0 reduces to reverse_kl, 1.0 to forward_kl.", + help=( + "Mixture coefficient for --opd-kl-type=jsd. Ordinary OPD uses alpha=0 for " + "KL(student||teacher) and alpha=1 for KL(teacher||student); SDPO uses the same " + "endpoint convention." + ), ) parser.add_argument( "--opd-norm-mode", @@ -630,10 +636,24 @@ def add_opd_arguments(parser: Any) -> Any: help=("Which token set the OPD KL signal is computed on."), ) parser.add_argument( - "--opd-teacher-prompt-key", + "--opd-feedback-class", type=str, default=None, - help=("Dataset field name that holds the teacher-side prompt for On-Policy Self-Distillation (OPSD). "), + help=( + "Fully qualified EnvironmentFeedback class binding an OPD-flavored algorithm " + "(OPSD / SDPO) to the shared rollout path. Defaults to " + "relax.utils.opd.feedback.OPDFeedback." + ), + ) + parser.add_argument( + "--opd-feedback-kwargs", + type=json.loads, + default=None, + help=( + "JSON string of constructor parameters for --opd-feedback-class, e.g. " + '\'{"teacher_prompt_key": "teacher_prompt"}\' for OPSDFeedback or ' + "'{\"success_reward_threshold\": 0.8}' for SDPOFeedback." + ), ) parser.add_argument( "--opd-teacher-image-key", @@ -687,6 +707,13 @@ def validate_opd_args(args: Namespace, *, is_sft: bool, log: Any = logger) -> No if not getattr(args, "use_opd", False): return + from relax.utils.opd.feedback import load_feedback + + load_feedback( + getattr(args, "opd_feedback_class", None), + getattr(args, "opd_feedback_kwargs", None), + ).validate_launch_args(args) + # OPD is enabled here. Backfill the routing key so teacher routing AND the # per-source metrics (compute_mopd_metrics) get a consistent value even when # the user didn't pass --opd-teacher-key. The arg default stays None so that @@ -704,6 +731,17 @@ def validate_opd_args(args: Namespace, *, is_sft: bool, log: Any = logger) -> No if token_selection != "student_sampled": if args.opd_log_prob_top_k <= 0: raise ValueError(f"--opd-token-selection={token_selection} requires --opd-log-prob-top-k > 0 ") + if getattr(args, "allgather_cp", False): + raise ValueError( + "OPD Top-K token selection is not compatible with --allgather-cp; " + "use the standard zig-zag context-parallel layout." + ) + if getattr(args, "context_parallel_size", 1) > 1 or getattr(args, "dynamic_context_parallel", False): + raise ValueError( + "OPD Top-K token selection is not compatible with context parallelism; " + "use --context-parallel-size 1 without --dynamic-context-parallel, " + "or --opd-token-selection=student_sampled." + ) kl_type = args.opd_kl_type if token_selection == "student_sampled" and kl_type not in ("reverse_kl", "low_var_kl"): @@ -726,14 +764,6 @@ def validate_opd_args(args: Namespace, *, is_sft: bool, log: Any = logger) -> No "--opd-kl-coef=0.0 --opd-loss-coef=X for loss mode." ) - if getattr(args, "opd_teacher_prompt_key", None) is not None: - if args.opd_type != "sglang": - raise ValueError( - "--opd-teacher-prompt-key currently only supports --opd-type=sglang " - f"(got --opd-type={args.opd_type}). The megatron teacher path does not " - "yet rebuild a teacher-side data_iterator from teacher_tokens." - ) - if getattr(args, "opd_teacher_image_key", None) is not None: if args.opd_type != "sglang": raise ValueError("--opd-teacher-image-key currently only supports --opd-type=sglang.") @@ -873,7 +903,7 @@ async def build_teacher_preexpanded_image_data(sample: "Sample") -> list | None: return None return [ { - "format": opd_opsd_worker.PREEXPANDED_RAW_FORMAT, + "format": PREEXPANDED_RAW_FORMAT, "images_b64": list(image_b64_list), "image_grid_thw": _to_jsonable(image_grid_thw), } @@ -897,7 +927,7 @@ async def build_student_preexpanded_image_data(sample: "Sample") -> list | None: cached = await _encode_images_b64(raw_images, "_student_image_b64_cache", student_mm_in) return [ { - "format": opd_opsd_worker.PREEXPANDED_RAW_FORMAT, + "format": PREEXPANDED_RAW_FORMAT, "images_b64": cached, "image_grid_thw": _to_jsonable(image_grid_thw), } @@ -913,7 +943,9 @@ def _get_opd_transfer_schema(args: Namespace) -> list[str]: def consume_opd_train_data(data_fields: list[str], args: Namespace) -> None: if not (getattr(args, "use_opd", False) and getattr(args, "opd_type", None) == "sglang"): return - data_fields.extend(_get_opd_transfer_schema(args)) + for field_name in _get_opd_transfer_schema(args): + if field_name not in data_fields: + data_fields.append(field_name) def consume_opd_advantage_data(data_fields: list[str], args: Namespace) -> None: @@ -1056,9 +1088,9 @@ def compute_log_probs_on_topk_token_ids( logsumexp = logits_max + exp_sum.log() # [S, 1] # === Step 2: vocab-parallel gather of logits at topk_token_ids === - from megatron.core import mpu - if process_group is not None: + from megatron.core import mpu + tp_world_size = mpu.get_tensor_model_parallel_world_size() tp_rank = mpu.get_tensor_model_parallel_rank() else: @@ -1069,6 +1101,12 @@ def compute_log_probs_on_topk_token_ids( vocab_end = vocab_start + vocab_size_per_rank ids = topk_token_ids.to(device=logits.device, dtype=torch.long) # [S, K] + global_vocab_size = vocab_size_per_rank * tp_world_size + if torch.any(ids >= global_vocab_size): + raise ValueError( + f"Top-K token ids must be below the global vocabulary size {global_vocab_size}; " + f"got max id {int(ids.max().detach().cpu())}." + ) invalid = ids < 0 safe_ids = ids.clamp(min=0) in_range = (safe_ids >= vocab_start) & (safe_ids < vocab_end) @@ -1117,6 +1155,9 @@ def compute_opd_kl( positions are set to ``-inf`` after clamp (so ``logsumexp`` ignores them) and their contributions are zeroed before ``.sum(dim=-1)``. ``None`` (topk path with fixed K) skips all masking — behavior unchanged. + + For ordinary OPD, JSD ``alpha=0`` and ``alpha=1`` are explicit endpoint + aliases for ``KL(student || teacher)`` and ``KL(teacher || student)``. """ s = student_log_probs.float() t = teacher_log_probs.float() @@ -1140,6 +1181,8 @@ def compute_opd_kl( if not (0.0 <= jsd_alpha <= 1.0): raise ValueError(f"jsd_alpha must be in [0, 1], got {jsd_alpha}") + distribution_mask = mask + def _add_tail(lp: torch.Tensor) -> torch.Tensor: log_s = torch.logsumexp(lp, dim=-1, keepdim=True).clamp(max=-1e-7) tail = torch.log(-torch.expm1(log_s)) @@ -1158,16 +1201,16 @@ def _norm(lp: torch.Tensor) -> torch.Tensor: s_t = s t_t = t - K_orig = mask.size(-1) if mask is not None else s_t.size(-1) + K_orig = distribution_mask.size(-1) if distribution_mask is not None else s_t.size(-1) if jsd_alpha == 0.0: result = s_t.exp() * (s_t - t_t) - if mask is not None: - result = result[..., :K_orig].masked_fill(~mask, 0.0) + if distribution_mask is not None: + result = result[..., :K_orig].masked_fill(~distribution_mask, 0.0) return result.sum(dim=-1) if jsd_alpha == 1.0: result = t_t.exp() * (t_t - s_t) - if mask is not None: - result = result[..., :K_orig].masked_fill(~mask, 0.0) + if distribution_mask is not None: + result = result[..., :K_orig].masked_fill(~distribution_mask, 0.0) return result.sum(dim=-1) log_1ma = torch.log(torch.tensor(1.0 - jsd_alpha, device=s_t.device, dtype=s_t.dtype)) @@ -1176,9 +1219,9 @@ def _norm(lp: torch.Tensor) -> torch.Tensor: kl_student = s_t.exp() * (s_t - m_t) kl_teacher = t_t.exp() * (t_t - m_t) - if mask is not None: - kl_student = kl_student[..., :K_orig].masked_fill(~mask, 0.0) - kl_teacher = kl_teacher[..., :K_orig].masked_fill(~mask, 0.0) + if distribution_mask is not None: + kl_student = kl_student[..., :K_orig].masked_fill(~distribution_mask, 0.0) + kl_teacher = kl_teacher[..., :K_orig].masked_fill(~distribution_mask, 0.0) return (1.0 - jsd_alpha) * kl_student.sum(dim=-1) + jsd_alpha * kl_teacher.sum(dim=-1) raise ValueError(f"Unknown opd_kl_type: {kl_type}. Choose one of {OPD_KL_TYPES}.") @@ -1197,6 +1240,12 @@ def compute_opd_kl_topk( ``mask``: optional bool ``[R, K]`` (union path only). ``None`` for fixed-K topk path — behavior unchanged. + + This is the ordinary OPD convention: ``reverse_kl`` maps to the student + expectation and ``forward_kl`` maps to the teacher expectation. The JSD + boundary values are explicit endpoint aliases: ``jsd_alpha=0`` is + ``KL(student || teacher)`` and ``jsd_alpha=1`` is + ``KL(teacher || student)``. SDPO reuses these same OPD equations. """ if kl_type in ("reverse_kl", "forward_kl", "jsd"): if kl_type == "reverse_kl": @@ -1349,16 +1398,21 @@ def apply_opd_to_advantages( advantages[i] = adv - args.opd_kl_coef * kl_term.detach() -def reduce_opd_loss(batch: RolloutBatch, values: torch.Tensor) -> torch.Tensor: +def reduce_opd_loss( + batch: RolloutBatch, + values: torch.Tensor, + loss_masks: list[torch.Tensor] | None = None, +) -> torch.Tensor: + loss_masks = batch["loss_masks"] if loss_masks is None else loss_masks chunks = torch.split(values, batch["response_lengths"], dim=0) masked_chunks = [] - for chunk, loss_mask in zip(chunks, batch["loss_masks"], strict=False): + for chunk, loss_mask in zip(chunks, loss_masks, strict=False): mask = loss_mask.to(device=chunk.device, dtype=chunk.dtype) masked = chunk * mask masked_chunks.append(masked) numerator = torch.cat(masked_chunks, dim=0).sum() - denominator = sum(mask.to(device=values.device, dtype=values.dtype).sum() for mask in batch["loss_masks"]) + denominator = sum(mask.to(device=values.device, dtype=values.dtype).sum() for mask in loss_masks) return numerator / torch.clamp_min(denominator, 1) @@ -1448,7 +1502,25 @@ def compute_policy_opd_loss( with torch.no_grad(): reported_loss["opd_is_clip_frac"] = (ratio > clip).float().mean().clone().detach() - opd_loss = reduce_opd_loss(batch, opd_per_token_kl) + if batch.get(OPD_SAMPLE_MASK) is not None: + # The batch loss_masks are already gated by the sample mask (get_batch), + # so the standard sum-style reduction excludes inactive samples from both + # the numerator and the global num_tokens normalizer. + from relax.backends.megatron.cp_utils import get_sum_of_sample_mean + + opd_loss = get_sum_of_sample_mean( + batch["total_lengths"], + batch["response_lengths"], + batch["loss_masks"], + getattr(args, "calculate_per_token_loss", False), + args.qkv_format, + batch.get("max_seq_lens"), + batch.get("padded_total_lengths"), + dynamic_cp_size=batch.get("dynamic_cp_size"), + dynamic_cp_rank=batch.get("dynamic_cp_rank"), + )(opd_per_token_kl) + else: + opd_loss = reduce_opd_loss(batch, opd_per_token_kl) return opd_loss_coef * opd_loss, reported_loss diff --git a/relax/utils/opd/opsd/__init__.py b/relax/utils/opd/opsd/__init__.py new file mode 100644 index 000000000..9f3863608 --- /dev/null +++ b/relax/utils/opd/opsd/__init__.py @@ -0,0 +1 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. diff --git a/relax/utils/opd/opsd/feedback.py b/relax/utils/opd/opsd/feedback.py new file mode 100644 index 000000000..67197dd20 --- /dev/null +++ b/relax/utils/opd/opsd/feedback.py @@ -0,0 +1,45 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""OPSD feedback: the dataset-provided teacher prompt is the privilege. + +The raw dataset column is rendered at ingestion and surfaced as +``metadata["opd_teacher_prompt"]``; this class owns the policy of assigning it +to ``sample.teacher_prompt``. Samples without the field fall back to the +student prompt via ``OpsdWorker``, matching plain OPD. +""" + +from __future__ import annotations + +from typing import Any + +from relax.utils.opd.feedback import EnvironmentFeedback +from relax.utils.types import Sample + + +class OPSDFeedback(EnvironmentFeedback): + def __init__(self, teacher_prompt_key: str | None = None, success_reward_threshold: float = 1.0) -> None: + if not teacher_prompt_key: + raise ValueError("OPSDFeedback requires {'teacher_prompt_key': ...} in --opd-feedback-kwargs.") + super().__init__(teacher_prompt_key=teacher_prompt_key, success_reward_threshold=success_reward_threshold) + + def prepare_teacher_prompts(self, group: list[Sample], rewards: list[Any]) -> None: + for sample in group: + sample.teacher_prompt = None + sample.opd_sample_mask = None + privileged = sample.metadata.get("opd_teacher_prompt") if isinstance(sample.metadata, dict) else None + if privileged is not None: + sample.teacher_prompt = privileged + + @classmethod + def validate_launch_args(cls, args: Any) -> None: + if args.opd_type != "sglang": + raise ValueError( + "OPSD prompt routing currently only supports --opd-type=sglang " + f"(got --opd-type={args.opd_type}). The megatron teacher path does not " + "yet rebuild a teacher-side data_iterator from teacher_tokens." + ) + + +__all__ = [ + "OPSDFeedback", +] diff --git a/relax/utils/opd/sdpo/__init__.py b/relax/utils/opd/sdpo/__init__.py new file mode 100644 index 000000000..f7f1bdded --- /dev/null +++ b/relax/utils/opd/sdpo/__init__.py @@ -0,0 +1,16 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +from relax.utils.opd.sdpo.constants import SDPO_TOKEN_SELECTION +from relax.utils.opd.sdpo.validation import ( + validate_sdpo_student_topk_ids, + validate_sdpo_text_only, + validate_sdpo_topk_payload, +) + + +__all__ = [ + "SDPO_TOKEN_SELECTION", + "validate_sdpo_text_only", + "validate_sdpo_student_topk_ids", + "validate_sdpo_topk_payload", +] diff --git a/relax/utils/opd/sdpo/constants.py b/relax/utils/opd/sdpo/constants.py new file mode 100644 index 000000000..f27d22efb --- /dev/null +++ b/relax/utils/opd/sdpo/constants.py @@ -0,0 +1,5 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +SDPO_TOKEN_SELECTION = "student_topk" + +__all__ = ["SDPO_TOKEN_SELECTION"] diff --git a/relax/utils/opd/sdpo/feedback.py b/relax/utils/opd/sdpo/feedback.py new file mode 100644 index 000000000..7e90c161f --- /dev/null +++ b/relax/utils/opd/sdpo/feedback.py @@ -0,0 +1,215 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""SDPO feedback: privileged teacher prompts built from group outcomes.""" + +from __future__ import annotations + +import copy +from collections import defaultdict +from typing import Any + +from relax.utils.opd.feedback import EnvironmentFeedback +from relax.utils.opd.opd_main_worker import TopkWorker +from relax.utils.opd.opd_opsd_worker import OpsdWorker +from relax.utils.opd.opd_utils import OPD_SAMPLE_MASK +from relax.utils.opd.sdpo.constants import SDPO_TOKEN_SELECTION +from relax.utils.opd.sdpo.validation import ( + validate_sdpo_student_topk_ids, + validate_sdpo_text_only, + validate_sdpo_topk_payload, +) +from relax.utils.types import Sample + + +def _has_sdpo_teacher_prompt(prompt: object) -> bool: + if isinstance(prompt, str): + return bool(prompt.strip()) + if isinstance(prompt, list): + return bool(prompt) + return False + + +def _clear_teacher_payload(sample: Sample) -> None: + """Drop outputs from an earlier teacher request without touching rollout + Top-K.""" + + for field_name in ( + "teacher_log_probs", + "teacher_topk_token_ids", + "teacher_topk_log_probs", + "teacher_at_student_topk_log_probs", + "student_at_teacher_topk_log_probs", + "opd_topk_token_ids", + "opd_topk_student_log_probs", + "opd_topk_teacher_log_probs", + "opd_topk_ksz", + "teacher_tokens", + "teacher_prompt_length", + ): + setattr(sample, field_name, None) + + +def _render_sdpo_teacher_prompt(sample: Sample, additions: list[str]) -> str | list[dict[str, str]]: + prompt = copy.deepcopy(sample.prompt) + suffix = "\n\n".join(additions + (["Now produce the best answer to the original problem."] if additions else [])) + if isinstance(prompt, list): + messages = prompt + if not messages or messages[-1].get("role") != "user": + messages.append({"role": "user", "content": ""}) + messages[-1]["content"] = f"{messages[-1].get('content', '')}\n\n{suffix}" + return messages + return f"{prompt}\n\n{suffix}" + + +def _set_sdpo_teacher_prompt(sample: Sample, additions: list[str]) -> None: + sample.teacher_prompt = ( + _render_sdpo_teacher_prompt(sample, additions) if additions else copy.deepcopy(sample.prompt) + ) + sample.opd_sample_mask = bool(additions) + sample.teacher_tokens = None + sample.teacher_prompt_length = None + + +def _is_successful_reward(reward: Any, threshold: float) -> bool: + value = reward.get("score", reward.get("reward")) if isinstance(reward, dict) else reward + try: + return float(value) >= threshold + except (TypeError, ValueError): + return False + + +def _is_format_or_truncation_error(sample: Sample, reward: Any) -> bool: + if isinstance(reward, dict) and reward.get("format_error"): + return True + status = getattr(sample, "status", None) + if status is not None: + return str(getattr(status, "value", status)).casefold() == "truncated" + return False + + +def _sdpo_group_key(sample: Sample, position: int) -> Any: + if sample.group_index is not None: + return ("group", sample.group_index) + metadata = sample.metadata if isinstance(sample.metadata, dict) else {} + uid = metadata.get("uid") + return ("uid", uid) if uid is not None else ("singleton", position) + + +def _prepare_sdpo_teacher_prompts(group: list[Sample], rewards: list[Any], threshold: float) -> None: + if len(group) != len(rewards): + raise ValueError(f"feedback requires one reward per sample: {len(group)} != {len(rewards)}") + by_group: dict[Any, list[Sample]] = defaultdict(list) + for position, sample in enumerate(group): + by_group[_sdpo_group_key(sample, position)].append(sample) + reward_by_id = {id(sample): reward for sample, reward in zip(group, rewards, strict=True)} + successful = { + key: [sample for sample in samples if _is_successful_reward(reward_by_id[id(sample)], threshold)] + for key, samples in by_group.items() + } + for key, samples in by_group.items(): + for sample in samples: + reward = reward_by_id[id(sample)] + is_success = _is_successful_reward(reward, threshold) + peer = next((candidate for candidate in successful[key] if candidate is not sample), None) + additions = [] + if is_success: + source = peer or sample + additions.append(f"\n{source.response}\n") + elif peer is not None: + additions.append(f"\n{peer.response}\n") + elif _is_format_or_truncation_error(sample, reward): + feedback = EnvironmentFeedback.feedback_text(sample, reward) + if feedback: + additions.append(f"\n{feedback}\n") + _set_sdpo_teacher_prompt(sample, additions) + + +class SDPOFeedback(EnvironmentFeedback): + def prepare_teacher_prompts(self, group: list[Sample], rewards: list[Any]) -> None: + for sample in group: + validate_sdpo_text_only(sample) + _prepare_sdpo_teacher_prompts(group, rewards, self.success_reward_threshold) + for sample in group: + _clear_teacher_payload(sample) + for sample in group: + if int(sample.response_length or 0) > 0 and not _has_sdpo_teacher_prompt(sample.teacher_prompt): + raise ValueError( + f"SDPO requires a teacher prompt for every non-empty response; sample_index={sample.index}" + ) + + @staticmethod + def create_opsd_worker(args: Any) -> OpsdWorker: + return OpsdWorker(is_opsd=True) + + @classmethod + def validate_launch_args(cls, args: Any) -> None: + if args.opd_type != "sglang": + raise ValueError("SDPO prompt routing requires --opd-type=sglang.") + if int(getattr(args, "pipeline_model_parallel_size", 1)) != 1: + raise ValueError("SDPO prompt routing does not support pipeline parallelism.") + if getattr(args, "enable_mtp_training", False): + raise ValueError("SDPO prompt routing does not support MTP auxiliary loss.") + if args.opd_token_selection != SDPO_TOKEN_SELECTION: + raise ValueError("SDPO prompt routing only supports --opd-token-selection=student_topk.") + if args.opd_kl_type not in ("forward_kl", "reverse_kl", "jsd"): + raise ValueError("SDPO prompt routing supports only forward_kl, reverse_kl, or jsd.") + if getattr(args, "opd_norm_mode", "tail") == "trunc": + raise ValueError("SDPO prompt routing requires --opd-norm-mode=tail or norm.") + if not getattr(args, "calculate_per_token_loss", False): + raise ValueError("SDPO prompt routing requires --calculate-per-token-loss.") + if getattr(args, "multimodal_keys", None) or any( + getattr(args, field_name, None) is not None + for field_name in ("opd_teacher_image_key", "opd_teacher_video_key", "opd_teacher_audio_key") + ): + raise ValueError("SDPO prompt routing only supports text inputs; multimodal fields are not supported.") + if not getattr(args, "group_rm", False): + raise ValueError("SDPO requires --group-rm to build privileged teacher prompts from group outcomes.") + opd_kl_coef = float(getattr(args, "opd_kl_coef", 0.0) or 0.0) + opd_loss_coef = float(getattr(args, "opd_loss_coef", 0.0) or 0.0) + if opd_kl_coef != 0.0 or opd_loss_coef <= 0.0: + raise ValueError( + "SDPO loss and prompt-routing teacher mode require --opd-kl-coef=0 and a positive --opd-loss-coef." + ) + + def extra_transfer_schema(self) -> list[str]: + return [OPD_SAMPLE_MASK] + + def produce_extra_transfer(self, samples: list[Sample], train_data: dict) -> None: + train_data[OPD_SAMPLE_MASK] = [bool(sample.opd_sample_mask) for sample in samples] + + def check_student_topk_ids(self, sample: Sample, top_k: int) -> None: + validate_sdpo_student_topk_ids( + token_ids=sample.student_topk_token_ids, + response_rows=int(sample.response_length or 0), + top_k=top_k, + sample_index=int(sample.index) if sample.index is not None else -1, + ) + + def check_transfer_channels(self, sample: Sample, channels: dict, top_k: int) -> None: + validate_sdpo_topk_payload( + token_ids=channels.get(TopkWorker.TRANSFER_TOKEN_IDS), + teacher_log_probs=channels.get(TopkWorker.TRANSFER_TEACHER_LOG_PROBS), + response_rows=int(sample.response_length or 0), + top_k=top_k, + sample_index=int(sample.index) if sample.index is not None else -1, + ) + + +class GoldenAnswerSDPOFeedback(SDPOFeedback): + """Static text datasets scored against a golden answer (MCQ, tool + calls).""" + + +class CodeSDPOFeedback(SDPOFeedback): + def record_sample_feedback(self, sample: Sample, reward: Any) -> None: + raise NotImplementedError("CodeSDPOFeedback is a placeholder; the code-domain reward is not wired up yet") + + def prepare_teacher_prompts(self, group: list[Sample], rewards: list[Any]) -> None: + raise NotImplementedError("CodeSDPOFeedback is a placeholder; the code-domain reward is not wired up yet") + + +__all__ = [ + "SDPOFeedback", + "GoldenAnswerSDPOFeedback", + "CodeSDPOFeedback", +] diff --git a/relax/utils/opd/sdpo/validation.py b/relax/utils/opd/sdpo/validation.py new file mode 100644 index 000000000..c18609bcd --- /dev/null +++ b/relax/utils/opd/sdpo/validation.py @@ -0,0 +1,90 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +from typing import Any + +import torch + + +def validate_sdpo_text_only(sample: Any) -> None: + """Reject multimodal or structured-content samples at the SDPO boundary.""" + + def present(value: Any) -> bool: + if value is None: + return False + if isinstance(value, dict): + return any(present(item) for item in value.values()) + if isinstance(value, (list, tuple, set)): + return any(present(item) for item in value) + if isinstance(value, str): + return bool(value) + if hasattr(value, "numel"): + return bool(value.numel()) + return True + + for field_name in ( + "multimodal_inputs", + "multimodal_train_inputs", + "teacher_multimodal_inputs", + "teacher_image_data", + "teacher_image_b64_list", + "teacher_image_grid_thw", + ): + value = getattr(sample, field_name, None) + if present(value): + raise ValueError(f"SDPO only supports text inputs; sample contains {field_name}") + + for field_name in ("prompt", "teacher_prompt"): + prompt = getattr(sample, field_name, None) + if prompt is None or isinstance(prompt, str): + continue + if not isinstance(prompt, list): + raise ValueError(f"SDPO only supports text {field_name}; expected a string or text messages") + for index, message in enumerate(prompt): + if not isinstance(message, dict) or not isinstance(message.get("role"), str): + raise ValueError(f"SDPO text message {index} in {field_name} must contain a string role") + if "content" in message and not isinstance(message["content"], str): + raise ValueError("Relax-SDPO only supports text chat message content") + + +def validate_sdpo_topk_payload( + *, token_ids: Any, teacher_log_probs: Any, response_rows: int, top_k: int, sample_index: int +) -> None: + if response_rows == 0: + return + if token_ids is None or teacher_log_probs is None: + raise ValueError(f"SDPO sample {sample_index} is missing its complete teacher Top-K payload") + expected_shape = (response_rows, top_k) + token_ids_tensor = torch.as_tensor(token_ids) + teacher_tensor = torch.as_tensor(teacher_log_probs) + if tuple(token_ids_tensor.shape) != expected_shape: + raise ValueError( + f"SDPO sample {sample_index} token-id payload shape {tuple(token_ids_tensor.shape)} " + f"does not match {expected_shape}" + ) + if tuple(teacher_tensor.shape) != expected_shape: + raise ValueError( + f"SDPO sample {sample_index} teacher payload shape {tuple(teacher_tensor.shape)} " + f"does not match {expected_shape}" + ) + if torch.isnan(teacher_tensor).any() or torch.isposinf(teacher_tensor).any(): + raise ValueError(f"SDPO sample {sample_index} teacher Top-K payload contains NaN or +inf") + + +def validate_sdpo_student_topk_ids(*, token_ids: Any, response_rows: int, top_k: int, sample_index: int) -> None: + if token_ids is None: + raise ValueError(f"SDPO sample {sample_index} is missing student Top-K token ids") + if response_rows <= 0 or top_k <= 0: + raise ValueError( + f"SDPO sample {sample_index} has invalid Top-K dimensions: rows={response_rows}, top_k={top_k}" + ) + token_ids_tensor = torch.as_tensor(token_ids) + expected_shape = (response_rows, top_k) + if tuple(token_ids_tensor.shape) != expected_shape: + raise ValueError( + f"SDPO sample {sample_index} student Top-K id shape {tuple(token_ids_tensor.shape)} " + f"does not match {expected_shape}" + ) + if token_ids_tensor.dtype == torch.bool or token_ids_tensor.dtype.is_floating_point: + raise ValueError(f"SDPO sample {sample_index} student Top-K ids must be integer token ids") + if token_ids_tensor.numel() and bool((token_ids_tensor < 0).any()): + raise ValueError(f"SDPO sample {sample_index} student Top-K ids contain a negative token id") diff --git a/relax/utils/types.py b/relax/utils/types.py index 9c7cabb5b..bc887fb37 100644 --- a/relax/utils/types.py +++ b/relax/utils/types.py @@ -43,6 +43,7 @@ class Sample: opd_topk_student_log_probs: list[np.ndarray] | None = None opd_topk_teacher_log_probs: list[np.ndarray] | None = None opd_topk_ksz: np.ndarray | None = None # [only union] + opd_sample_mask: bool | None = None teacher_prompt: str | list[dict[str, str]] | None = None teacher_multimodal_inputs: dict[str, Any] | None = None @@ -212,7 +213,7 @@ class ParamInfo: # A dict-based batch produced along the rollout -> training path # In Megatron backend, several fields are converted to torch.Tensor lists on GPU # before being consumed by data iterators (see megatron_utils.actor._get_rollout_data). -RolloutBatch = dict[str, list[torch.Tensor] | list[int] | list[float] | list[str]] +RolloutBatch = dict[str, list[torch.Tensor] | list[int] | list[float] | list[bool] | list[str]] SFTBatch = dict[str, list[torch.Tensor] | list[int] | list[str] | dict | None] """SFT 训练 batch 的 dict alias。 diff --git a/relax/utils/utils.py b/relax/utils/utils.py index bac105d6c..f398b62b2 100644 --- a/relax/utils/utils.py +++ b/relax/utils/utils.py @@ -248,10 +248,12 @@ def dict_to_tensordict( if not data: return TensorDict({}, batch_size=0 if batch_size is None else batch_size, device=device) - def _nesting_depth(x): - if isinstance(x, list) and x: - return 1 + _nesting_depth(x[0]) - return 0 + def _nesting_depth(x, *, root: bool = True): + if not isinstance(x, list): + return 0 + if not x: + return 0 if root else 1 + return 1 + max(_nesting_depth(item, root=False) for item in x) def _scalar_dtype(sample) -> Optional[torch.dtype]: """Return an explicit dtype only for bool/float; None lets torch.tensor @@ -267,8 +269,23 @@ def _to_tensor_1d(lst): dtype = _scalar_dtype(lst[0]) return torch.tensor(lst, dtype=dtype, device=device) + def _first_scalar(value): + if isinstance(value, list): + for item in value: + scalar = _first_scalar(item) + if scalar is not None: + return scalar + return None + return value + def _to_tensor_2d(lst): - dtype = _scalar_dtype(lst[0][0]) + nonempty_row = next((row for row in lst if row), None) + if nonempty_row is None: + dtype = None + else: + dtype = _scalar_dtype(_first_scalar(nonempty_row)) + if dtype is None: + dtype = torch.as_tensor(nonempty_row).dtype tensors = [torch.tensor(seq, dtype=dtype, device=device) for seq in lst] return torch.nested.as_nested_tensor(tensors, layout=torch.jagged) diff --git a/scripts/models/qwen3-4B-Instruct-2507.sh b/scripts/models/qwen3-4B-Instruct-2507.sh new file mode 100644 index 000000000..d84937e33 --- /dev/null +++ b/scripts/models/qwen3-4B-Instruct-2507.sh @@ -0,0 +1,19 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +MODEL_ARGS=( + --swiglu + --num-layers 36 + --hidden-size 2560 + --ffn-hidden-size 9728 + --num-attention-heads 32 + --group-query-attention + --num-query-groups 8 + --use-rotary-position-embeddings + --disable-bias-linear + --normalization "RMSNorm" + --norm-epsilon 1e-6 + --rotary-base "${MODEL_ARGS_ROTARY_BASE:-5000000}" + --vocab-size 151936 + --kv-channels 128 + --qk-layernorm +) diff --git a/tests/backends/megatron/test_data_vpp.py b/tests/backends/megatron/test_data_vpp.py index c2f8980aa..2de455934 100644 --- a/tests/backends/megatron/test_data_vpp.py +++ b/tests/backends/megatron/test_data_vpp.py @@ -182,3 +182,57 @@ def test_get_data_iterator_balance_data_without_boundaries_uses_regular_steps(mo _, num_microbatches = data_module.get_data_iterator(args, object(), rollout_data) assert num_microbatches == [2, 2] + + +def test_opd_sample_mask_gates_loss_masks_per_sample(monkeypatch): + data_module = _load_data_module(monkeypatch) + batch = { + "loss_masks": [torch.ones(2), torch.ones(1), torch.ones(3), torch.ones(2)], + data_module.OPD_SAMPLE_MASK: [True, False, True, False], + } + + data_module._apply_opd_sample_mask(batch) + + assert torch.equal(batch["loss_masks"][0], torch.ones(2)) + assert torch.equal(batch["loss_masks"][1], torch.zeros(1)) + assert torch.equal(batch["loss_masks"][2], torch.ones(3)) + assert torch.equal(batch["loss_masks"][3], torch.zeros(2)) + + +def test_opd_sample_mask_gates_loss_masks_after_microbatch_permutation(monkeypatch): + data_module = _load_data_module(monkeypatch) + rollout_data = { + "total_lengths": [4, 3, 6, 5], + "loss_masks": [torch.ones(2), torch.ones(1), torch.ones(3), torch.ones(2)], + data_module.OPD_SAMPLE_MASK: [True, False, True, False], + } + iterator = data_module.DataIterator(rollout_data, micro_batch_indices=[[1], [0], [3], [2]]) + + for _ in range(4): + batch = iterator.get_next(["total_lengths", "loss_masks", data_module.OPD_SAMPLE_MASK]) + data_module._apply_opd_sample_mask(batch) + for loss_mask, sample_mask in zip(batch["loss_masks"], batch[data_module.OPD_SAMPLE_MASK], strict=True): + assert torch.equal( + loss_mask, torch.zeros(loss_mask.numel()) if not sample_mask else torch.ones_like(loss_mask) + ) + + +def test_opd_sample_mask_leaves_ordinary_batch_untouched(monkeypatch): + data_module = _load_data_module(monkeypatch) + batch = {"loss_masks": [torch.ones(2), torch.ones(1)]} + + data_module._apply_opd_sample_mask(batch) + + assert torch.equal(batch["loss_masks"][0], torch.ones(2)) + assert torch.equal(batch["loss_masks"][1], torch.ones(1)) + + +def test_opd_sample_mask_rejects_length_mismatch(monkeypatch): + data_module = _load_data_module(monkeypatch) + batch = { + "loss_masks": [torch.ones(2), torch.ones(1)], + data_module.OPD_SAMPLE_MASK: [True], + } + + with pytest.raises(ValueError, match="one scalar per sample"): + data_module._apply_opd_sample_mask(batch) diff --git a/tests/engine/rollout/test_on_policy_distillation_payload.py b/tests/engine/rollout/test_on_policy_distillation_payload.py index dbd99f61e..96e2a82bf 100644 --- a/tests/engine/rollout/test_on_policy_distillation_payload.py +++ b/tests/engine/rollout/test_on_policy_distillation_payload.py @@ -9,10 +9,19 @@ carried as sglang base64 fields, decoded into numpy arrays. """ +import asyncio +import base64 as pybase64 +from types import SimpleNamespace + import numpy as np -import pybase64 +import pytest +from relax.engine.rollout.on_policy_distillation import OpdManager +from relax.utils.opd.feedback import OPDFeedback from relax.utils.opd.opd_main_worker import LogprobResponse, TopkWorker +from relax.utils.opd.opd_opsd_worker import OpsdWorker +from relax.utils.opd.sdpo.feedback import SDPOFeedback, _clear_teacher_payload +from relax.utils.types import Sample def _b64(arr: np.ndarray) -> str: @@ -48,6 +57,37 @@ def test_base_logprobs_1d_from_b64_and_legacy() -> None: np.testing.assert_allclose(LogprobResponse(r2).base_logprobs_1d(), np.array([-0.5, -0.7], dtype=np.float32)) +def test_plain_sglang_topk_fields_are_decoded() -> None: + response = { + "meta_info": { + "output_top_logprobs": [ + [(-0.1, 0, None), (-0.2, 5, None)], + [(-0.3, 7, None), (-0.4, 9, None)], + ], + "input_top_logprobs": [ + [(-0.5, 1, None), (-0.6, 2, None)], + [(-0.7, 3, None), (-0.8, 4, None)], + [(-0.9, 6, None), (-1.0, 8, None)], + ], + "input_token_ids_logprobs": [ + [(-0.5, 1, None), (-0.6, 2, None)], + [(-0.7, 3, None), (-0.8, 4, None)], + [(-0.9, 6, None), (-1.0, 8, None)], + ], + } + } + + output_ids, output_lps = LogprobResponse(response).self_topk("rollout", top_k=2) + input_ids, input_lps = LogprobResponse(response).self_topk("prefill", top_k=2, response_length=1) + input_query_lps = LogprobResponse(response).other_topk(response_length=1, top_k=2) + + np.testing.assert_array_equal(output_ids, np.array([[0, 5], [7, 9]], dtype=np.int32)) + np.testing.assert_allclose(output_lps, [[-0.1, -0.2], [-0.3, -0.4]]) + np.testing.assert_array_equal(input_ids, np.array([[6, 8]], dtype=np.int32)) + np.testing.assert_allclose(input_lps, [[-0.9, -1.0]]) + np.testing.assert_allclose(input_query_lps, [[-0.9, -1.0]]) + + def test_build_teacher_payload_adds_student_topk_query_ids() -> None: # student_topk + kl_coef != 0 => teacher_at_student=True => token_ids_logprob present. w = TopkWorker("student_topk", top_k=2, opd_kl_coef=1.0, opd_loss_coef=0.0) @@ -67,6 +107,324 @@ def test_build_teacher_payload_adds_student_topk_query_ids() -> None: assert payload["token_ids_logprob"] == [3, 5, 7, 9, 0, 0] +def test_rollout_topk_payload_parsing_keeps_ordinary_opd_contract() -> None: + manager = object.__new__(OpdManager) + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=0.0) + meta_info = { + "output_token_logprobs_val_b64": _b64(np.array([-0.1, -0.2], dtype=np.float32)), + "output_token_logprobs_idx_b64": _b64(np.array([0, 5], dtype=np.int32)), + } + + tokens, log_probs = manager.parse_rollout_logprobs(meta_info, [9, 9], [-1.0, -1.0]) + + assert tokens == [0, 5] + np.testing.assert_allclose(log_probs, [-0.1, -0.2]) + + +def test_sdpo_teacher_prefill_uses_privileged_prompt_offset(monkeypatch) -> None: + manager = object.__new__(OpdManager) + manager.args = type("Args", (), {"opd_teacher_url": "http://teacher"})() + manager.feedback = SDPOFeedback() + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=1.0) + manager.sampled_worker = None + manager.opsd_worker = OpsdWorker(is_opsd=True) + + sample = Sample( + tokens=[10, 11, 20, 21], + response_length=2, + teacher_prompt=[{"role": "user", "content": "privileged context"}], + teacher_tokens=[100, 101, 102, 20, 21], + teacher_prompt_length=3, + student_topk_token_ids=np.array([[7, 8], [9, 10]], dtype=np.int32), + ) + captured = {} + + async def fake_post_logprob(session, url, payload, requested_sample, err_tag): + captured.update(url=url, payload=payload, sample=requested_sample, err_tag=err_tag) + return LogprobResponse( + { + "meta_info": { + "input_token_logprobs_val_b64": _b64(np.array([-0.1, -0.2, -0.3], dtype=np.float32)), + "input_token_ids_logprobs_val_b64": _b64(np.array([-0.4, -0.5, -0.6, -0.7], dtype=np.float32)), + } + } + ) + + monkeypatch.setattr(manager, "_post_logprob", fake_post_logprob) + + assert asyncio.run(manager._teacher_prefill(sample, object())) is True + assert captured["url"] == "http://teacher" + assert captured["payload"]["input_ids"] == sample.teacher_tokens + assert captured["payload"]["logprob_start_len"] == 2 + assert captured["payload"]["token_ids_logprob"] == [7, 8, 9, 10, 0, 0] + assert sample.teacher_at_student_topk_log_probs.shape == (2, 2) + + +def test_ordinary_teacher_prefill_keeps_rollout_input_and_original_offset(monkeypatch) -> None: + manager = object.__new__(OpdManager) + manager.args = type("Args", (), {"opd_teacher_url": "http://teacher"})() + manager.feedback = OPDFeedback() + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=1.0) + manager.sampled_worker = None + manager.opsd_worker = OpsdWorker(is_opsd=True) + + sample = Sample( + tokens=[10, 11, 20, 21], + rollout_tokens=[1, 2, 3, 4, 20, 21], + response_length=2, + student_topk_token_ids=np.array([[7, 8], [9, 10]], dtype=np.int32), + ) + captured = {} + + async def fake_post_logprob(session, url, payload, requested_sample, err_tag): + captured.update(url=url, payload=payload, sample=requested_sample, err_tag=err_tag) + return LogprobResponse( + { + "meta_info": { + "input_token_logprobs_val_b64": _b64(np.array([-0.1, -0.2, -0.3], dtype=np.float32)), + "input_token_ids_logprobs_val_b64": _b64(np.array([-0.4, -0.5, -0.6, -0.7], dtype=np.float32)), + } + } + ) + + monkeypatch.setattr(manager, "_post_logprob", fake_post_logprob) + + assert asyncio.run(manager._teacher_prefill(sample, object())) is True + assert captured["payload"]["input_ids"] == sample.rollout_tokens + assert captured["payload"]["logprob_start_len"] == 1 + assert captured["payload"]["token_ids_logprob"] == [7, 8, 9, 10, 0, 0] + assert "image_data" not in captured["payload"] + + +@pytest.mark.parametrize( + "student_topk_token_ids", + [ + None, + np.array([[7, 8]], dtype=np.int32), + np.array([[7, -1], [9, 10]], dtype=np.int32), + ], +) +def test_sdpo_rejects_invalid_student_topk_before_teacher_request( + monkeypatch, + student_topk_token_ids, +) -> None: + manager = object.__new__(OpdManager) + manager.args = type("Args", (), {"opd_teacher_url": "http://teacher"})() + manager.feedback = SDPOFeedback() + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=1.0) + manager.sampled_worker = None + manager.opsd_worker = OpsdWorker(is_opsd=True) + sample = Sample( + index=8, + tokens=[10, 11, 20, 21], + response_length=2, + teacher_prompt="privileged context", + student_topk_token_ids=student_topk_token_ids, + ) + + async def fail_if_requested(*args, **kwargs): + raise AssertionError("SDPO must validate rollout Top-K ids before the teacher request") + + monkeypatch.setattr(manager, "_post_logprob", fail_if_requested) + + with pytest.raises(ValueError, match="student Top-K"): + asyncio.run(manager._teacher_prefill(sample, object())) + + +def test_sdpo_manager_constructs_opsd_worker_and_feedback() -> None: + args = type( + "Args", + (), + { + "opd_token_selection": "student_topk", + "opd_log_prob_top_k": 2, + "opd_kl_coef": 0.0, + "opd_loss_coef": 1.0, + "opd_feedback_class": "relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback", + "opd_feedback_kwargs": None, + "opd_teacher_image_key": None, + }, + )() + + manager = OpdManager(args) + + assert manager.opsd_worker is not None + assert manager.opsd_worker.is_opsd + assert isinstance(manager.feedback, SDPOFeedback) + + +def test_sdpo_transfer_schema_contains_sample_mask_and_topk_payload() -> None: + manager = object.__new__(OpdManager) + manager.feedback = SDPOFeedback() + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=1.0) + manager.sampled_worker = None + + assert manager.schema_opd_transfer_data() == [ + "opd_sample_mask", + TopkWorker.TRANSFER_TOKEN_IDS, + TopkWorker.TRANSFER_TEACHER_LOG_PROBS, + ] + + +def test_sdpo_transfer_preserves_sample_mask_order() -> None: + manager = object.__new__(OpdManager) + manager.feedback = SDPOFeedback() + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=1.0) + manager.sampled_worker = None + samples = [Sample(index=0, opd_sample_mask=True), Sample(index=1, opd_sample_mask=False)] + train_data = {} + + manager.produce_opd_transfer_data(samples, train_data) + + assert train_data["opd_sample_mask"] == [True, False] + + +@pytest.mark.parametrize( + ("mode", "selection", "kl_coef", "loss_coef", "expected"), + [ + ( + "opd", + "student_sampled", + 0.0, + 1.0, + ["teacher_log_probs", "rollout_log_probs"], + ), + ("opd", "student_topk", 0.0, 1.0, ["opd_topk_token_ids", "opd_topk_teacher_log_probs"]), + ( + "opd", + "student_topk", + 1.0, + 0.0, + ["opd_topk_token_ids", "opd_topk_teacher_log_probs", "opd_topk_student_log_probs"], + ), + ("opd", "teacher_topk", 0.0, 1.0, ["opd_topk_token_ids", "opd_topk_teacher_log_probs"]), + ("opd", "union", 0.0, 1.0, ["opd_topk_token_ids", "opd_topk_teacher_log_probs", "opd_topk_ksz"]), + ("opsd", "student_topk", 0.0, 1.0, ["opd_topk_token_ids", "opd_topk_teacher_log_probs"]), + ( + "sdpo", + "student_topk", + 0.0, + 1.0, + ["opd_sample_mask", "opd_topk_token_ids", "opd_topk_teacher_log_probs"], + ), + ], +) +def test_teacher_transfer_schema_matrix(mode, selection, kl_coef, loss_coef, expected) -> None: + args = SimpleNamespace( + use_opd=True, + opd_type="sglang", + group_rm=mode == "sdpo", + opd_feedback_class="relax.utils.opd.sdpo.feedback.GoldenAnswerSDPOFeedback" if mode == "sdpo" else None, + opd_feedback_kwargs={"teacher_prompt_key": "teacher_prompt"} if mode == "opsd" else None, + opd_token_selection=selection, + opd_log_prob_top_k=2, + opd_kl_coef=kl_coef, + opd_loss_coef=loss_coef, + opd_teacher_image_key=None, + ) + + schema = OpdManager(args).schema_opd_transfer_data() + + assert schema == expected + + +def test_sdpo_assembly_rejects_missing_teacher_topk_payload() -> None: + manager = object.__new__(OpdManager) + manager.feedback = SDPOFeedback() + manager.topk_worker = TopkWorker("student_topk", top_k=2, opd_loss_coef=1.0) + sample = Sample( + index=4, + response_length=1, + student_topk_token_ids=np.array([[1, 2]], dtype=np.int32), + student_topk_log_probs=np.array([[-0.1, -0.2]], dtype=np.float32), + ) + + with pytest.raises(ValueError, match="complete teacher Top-K payload"): + manager._assemble_transfer([sample]) + + +def test_teacher_payload_reset_preserves_student_rollout_topk() -> None: + sample = Sample( + student_topk_token_ids=np.array([[1, 2]], dtype=np.int32), + student_topk_log_probs=np.array([[-0.1, -0.2]], dtype=np.float32), + teacher_log_probs=[-0.3], + teacher_topk_token_ids=np.array([[3, 4]], dtype=np.int32), + teacher_topk_log_probs=np.array([[-0.3, -0.4]], dtype=np.float32), + teacher_at_student_topk_log_probs=np.array([[-0.5, -0.6]], dtype=np.float32), + student_at_teacher_topk_log_probs=np.array([[-0.7, -0.8]], dtype=np.float32), + opd_topk_token_ids=[np.array([[3, 4]], dtype=np.int32)], + opd_topk_teacher_log_probs=[np.array([[-0.3, -0.4]], dtype=np.float32)], + teacher_tokens=[9, 10], + teacher_prompt_length=1, + ) + + _clear_teacher_payload(sample) + + assert sample.teacher_log_probs is None + assert sample.teacher_topk_token_ids is None + assert sample.teacher_topk_log_probs is None + assert sample.teacher_at_student_topk_log_probs is None + assert sample.student_at_teacher_topk_log_probs is None + assert sample.opd_topk_token_ids is None + assert sample.opd_topk_teacher_log_probs is None + assert sample.teacher_tokens is None + assert sample.teacher_prompt_length is None + np.testing.assert_array_equal(sample.student_topk_token_ids, np.array([[1, 2]], dtype=np.int32)) + np.testing.assert_array_equal(sample.student_topk_log_probs, np.array([[-0.1, -0.2]], dtype=np.float32)) + + +def test_ordinary_prefill_does_not_clear_existing_teacher_payload(monkeypatch) -> None: + manager = object.__new__(OpdManager) + manager.args = type( + "Args", + (), + {"opd_teacher_connector_limit": 1, "opd_teacher_timeout_s": 1.0}, + )() + manager.feedback = OPDFeedback() + manager.topk_worker = None + manager.sampled_worker = None + manager.opsd_worker = OpsdWorker(is_opsd=True) + sample = Sample(response_length=1, teacher_log_probs=[-0.25]) + + async def fake_teacher_prefill(requested_sample, session): + return True + + monkeypatch.setattr(manager, "_teacher_prefill", fake_teacher_prefill) + monkeypatch.setattr(manager, "_assemble_transfer", lambda samples: None) + monkeypatch.setattr(manager, "_raise_if_all_failed", lambda samples, results: None) + + asyncio.run(manager.prefill(sample)) + + assert sample.teacher_log_probs == [-0.25] + + +def test_ordinary_prefill_keeps_zero_response_teacher_dispatch(monkeypatch) -> None: + manager = object.__new__(OpdManager) + manager.args = type( + "Args", + (), + {"opd_teacher_connector_limit": 1, "opd_teacher_timeout_s": 1.0}, + )() + manager.feedback = OPDFeedback() + manager.topk_worker = None + manager.sampled_worker = None + manager.opsd_worker = OpsdWorker(is_opsd=True) + sample = Sample(response_length=0) + requested = [] + + async def fake_teacher_prefill(requested_sample, session): + requested.append(requested_sample) + return True + + monkeypatch.setattr(manager, "_teacher_prefill", fake_teacher_prefill) + monkeypatch.setattr(manager, "_assemble_transfer", lambda samples: None) + monkeypatch.setattr(manager, "_raise_if_all_failed", lambda samples, results: None) + + asyncio.run(manager.prefill(sample)) + + assert requested == [sample] + + def test_build_transfer_channels_union_merges_and_pads() -> None: # union + kl_coef != 0 (is_advantage=True): merge student/teacher self-topk # ids per position, keep first-seen logprobs, pad ragged rows. diff --git a/tests/examples/sdpo/test_data_reward.py b/tests/examples/sdpo/test_data_reward.py new file mode 100644 index 000000000..f692b7a55 --- /dev/null +++ b/tests/examples/sdpo/test_data_reward.py @@ -0,0 +1,286 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +from types import SimpleNamespace + +import pytest + +from examples.on_policy_distillation.sdpo.prepare_data import normalize_rows +from examples.on_policy_distillation.sdpo.reward import score + + +def _sample(metadata: dict, response: str) -> SimpleNamespace: + return SimpleNamespace(metadata=metadata, response=response, label="") + + +def test_normalize_sciknoweval_filters_l3_domains_and_preserves_split() -> None: + rows = [ + { + "question": "Which option?", + "choices": {"text": ["one", "two"], "label": ["A", "B"]}, + "answerKey": "B", + "type": "mcq-2-choices", + "domain": "Physics", + "details": {"level": "L3"}, + }, + { + "question": "filtered", + "choices": {"text": [], "label": []}, + "answerKey": "A", + "type": "mcq-2-choices", + "domain": "Physics", + "details": {"level": "L2"}, + }, + ] + + normalized = normalize_rows("sciknoweval", rows, source_split="train") + + assert len(normalized) == 1 + assert normalized[0]["metadata"]["source_split"] == "train" + assert normalized[0]["metadata"]["domain"] == "Physics" + + +def test_normalize_rows_rejects_unknown_split() -> None: + with pytest.raises(ValueError, match="expected 'train' or 'test'"): + normalize_rows("sciknoweval", [], source_split="validation") + + +def test_normalize_sciknoweval_accepts_reference_flat_domain_format() -> None: + rows = [ + { + "idx": 7, + "dataset": "sciknoweval", + "kind": "mcq", + "answer": "C", + "prompt": "Question\n\nA: one\nB: two\nC: three\nD: four", + "system": "Return only the answer tag.", + } + ] + + normalized = normalize_rows("sciknoweval", rows, source_split="train", domain="material") + + assert len(normalized) == 1 + assert normalized[0]["label"] == "C" + assert normalized[0]["metadata"]["domain"] == "Materials" + assert "Return only the answer tag." in normalized[0]["prompt"] + + +def test_toolalpaca_reward_accepts_canonical_action_input() -> None: + sample = _sample( + { + "data_source": "toolalpaca", + "golden_answer": [{"Action": "search", "Action_Input": '{"query": "relax"}'}], + }, + 'Action: search\nAction Input: {"query": "relax"}', + ) + + result = score(None, sample) + + assert result["score"] == 1.0 + assert result["feedback"] == "" + + +def test_toolalpaca_reward_parses_nested_action_input_json() -> None: + sample = _sample( + { + "data_source": "toolalpaca", + "golden_answer": [ + { + "Action": "sendHttpRequest", + "Action_Input": ( + '{"method": "POST", "data": {"name": "John Doe", "email": "john.doe@example.com"}}' + ), + } + ], + }, + ( + "Action: sendHttpRequest\nAction Input: " + '{"method": "POST", "data": ' + '{"name": "John Doe", "email": "john.doe@example.com"}}' + ), + ) + + result = score(None, sample) + + assert result["score"] == 1.0 + assert result["feedback"] == "" + + +def test_toolalpaca_reward_accepts_action_names_with_spaces_and_brackets() -> None: + for action in ("Get Task", "Search Twitch", "[Optional]updateUserProfile"): + sample = _sample( + { + "data_source": "toolalpaca", + "golden_answer": [{"Action": action, "Action_Input": "{}"}], + }, + f"Action: {action}\nAction Input: {{}}", + ) + + assert score(None, sample)["score"] == 1.0 + + +def test_reference_tooluse_row_is_normalized_for_the_same_reward() -> None: + rows = [ + { + "idx": 3, + "dataset": "tooluse", + "kind": "tooluse", + "prompt": "Action: \nAction Input: ", + "answer": '[{"Action": "search", "Action_Input": "{\\"query\\": \\"relax\\"}"}]', + } + ] + + normalized = normalize_rows("tooluse", rows, source_split="train") + + assert len(normalized) == 1 + assert normalized[0]["metadata"]["data_source"] == "tooluse" + assert normalized[0]["metadata"]["golden_answer"][0]["Action"] == "search" + + +def test_toolalpaca_reward_wrong_action_has_no_feedback() -> None: + sample = _sample( + { + "data_source": "toolalpaca", + "golden_answer": [{"Action": "search", "Action_Input": '{"query": "relax"}'}], + }, + 'Action: lookup\nAction Input: {"query": "relax"}', + ) + + result = score(None, sample) + + assert result["score"] == 0.0 + assert result["format_error"] == 0 + assert result["feedback"] == "" + assert "search" not in result["feedback"] + + +_MULTI_STEP_GOLDEN = [ + {"Action": "getClientRequestData", "Action_Input": '{"url": "https://httpbin.org/get"}'}, + {"Action": "sendHttpRequest", "Action_Input": '{"method": "POST", "url": "https://httpbin.org/post"}'}, +] + + +def _multi_step_response() -> str: + return ( + "Action: getClientRequestData\n" + 'Action Input: {"url": "https://httpbin.org/get"}\n' + "Action: sendHttpRequest\n" + 'Action Input: {"method": "POST", "url": "https://httpbin.org/post"}' + ) + + +def test_tooluse_reward_accepts_correct_multi_step_trajectory() -> None: + sample = _sample({"data_source": "tooluse", "golden_answer": _MULTI_STEP_GOLDEN}, _multi_step_response()) + + result = score(None, sample) + + assert result["score"] == 1.0 + assert result["format_error"] == 0 + + +def test_tooluse_reward_rejects_swapped_step_order() -> None: + swapped = ( + "Action: sendHttpRequest\n" + 'Action Input: {"method": "POST", "url": "https://httpbin.org/post"}\n' + "Action: getClientRequestData\n" + 'Action Input: {"url": "https://httpbin.org/get"}' + ) + sample = _sample({"data_source": "tooluse", "golden_answer": _MULTI_STEP_GOLDEN}, swapped) + + result = score(None, sample) + + assert result["score"] == 0.0 + + +def test_tooluse_reward_rejects_cross_step_value_swap() -> None: + swapped_values = ( + "Action: getClientRequestData\n" + 'Action Input: {"method": "POST", "url": "https://httpbin.org/post"}\n' + "Action: sendHttpRequest\n" + 'Action Input: {"url": "https://httpbin.org/get"}' + ) + sample = _sample({"data_source": "tooluse", "golden_answer": _MULTI_STEP_GOLDEN}, swapped_values) + + result = score(None, sample) + + assert result["score"] == 0.0 + + +def _sciknoweval_sample(response: str, *, expected: str = "B", task_type: str = "mcq") -> SimpleNamespace: + return _sample( + { + "data_source": "sciknoweval", + "answer_key": expected, + "task_type": task_type, + }, + response, + ) + + +def test_sciknoweval_reward_accepts_answer_tag() -> None: + result = score(None, _sciknoweval_sample("The answer is B.")) + + assert result["score"] == 1.0 + assert result["format_error"] == 0 + assert result["feedback"] == "" + + +def test_sciknoweval_reward_missing_answer_tag_is_format_error() -> None: + result = score(None, _sciknoweval_sample("The answer is B.")) + + assert result["score"] == 0.0 + assert result["format_error"] == 1 + assert "wrong format" in result["feedback"] + + +def test_sciknoweval_reward_wrong_answer_has_no_feedback() -> None: + result = score(None, _sciknoweval_sample("The answer is C.")) + + assert result["score"] == 0.0 + assert result["format_error"] == 0 + assert result["feedback"] == "" + + +def test_sciknoweval_reward_truncation_feedback_takes_priority() -> None: + from relax.utils.types import Sample + + sample = _sciknoweval_sample("The answer is B.") + sample.status = Sample.Status.TRUNCATED + + result = score(None, sample) + + assert result["score"] == 0.0 + assert "truncated" in result["feedback"] + assert "wrong format" not in result["feedback"] + + +def test_toolalpaca_reward_missing_format_has_feedback() -> None: + sample = _sample( + { + "data_source": "toolalpaca", + "golden_answer": [{"Action": "search", "Action_Input": '{"query": "relax"}'}], + }, + "just a text answer", + ) + + result = score(None, sample) + + assert result["score"] == 0.0 + assert result["format_error"] == 1 + assert "format" in result["feedback"] + + +def test_toolalpaca_reward_truncation_has_feedback() -> None: + from relax.utils.types import Sample + + sample = _sample( + { + "data_source": "toolalpaca", + "golden_answer": [{"Action": "search", "Action_Input": '{"query": "relax"}'}], + }, + 'Action: lookup\nAction Input: {"query": "relax"}', + ) + sample.status = Sample.Status.TRUNCATED + + result = score(None, sample) + + assert "truncated" in result["feedback"] diff --git a/tests/examples/sdpo/test_prepare_data.py b/tests/examples/sdpo/test_prepare_data.py new file mode 100644 index 000000000..a373c4933 --- /dev/null +++ b/tests/examples/sdpo/test_prepare_data.py @@ -0,0 +1,177 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +import json +import sys + +from examples.on_policy_distillation.sdpo.prepare_data import _read_rows, main + + +def test_read_rows_accepts_jsonl_with_json_suffix(tmp_path) -> None: + path = tmp_path / "train.json" + path.write_text( + "\n".join(json.dumps(row) for row in ({"idx": 1}, {"idx": 2})) + "\n", + encoding="utf-8", + ) + + assert _read_rows(path) == [{"idx": 1}, {"idx": 2}] + + +def test_read_rows_accepts_json_array(tmp_path) -> None: + path = tmp_path / "train.json" + path.write_text(json.dumps([{"idx": 1}, {"idx": 2}]), encoding="utf-8") + + assert _read_rows(path) == [{"idx": 1}, {"idx": 2}] + + +def test_prepare_data_main_writes_normalized_jsonl(tmp_path, monkeypatch) -> None: + input_path = tmp_path / "chemistry" / "train.json" + output_path = tmp_path / "prepared" / "chemistry.jsonl" + input_path.parent.mkdir() + input_path.write_text( + json.dumps( + { + "idx": 4, + "dataset": "sciknoweval", + "kind": "mcq", + "answer": "B", + "prompt": "Choose B.", + "system": "Return only the answer.", + } + ) + + "\n", + encoding="utf-8", + ) + monkeypatch.setattr( + sys, + "argv", + [ + "prepare_data", + "--dataset", + "sciknoweval", + "--input", + str(input_path), + "--domain", + "chemistry", + "--source-split", + "train", + "--output", + str(output_path), + ], + ) + + main() + + row = json.loads(output_path.read_text(encoding="utf-8").strip()) + assert row["label"] == "B" + assert row["metadata"]["source_split"] == "train" + assert row["metadata"]["domain"] == "Chemistry" + + +def test_prepare_data_main_limits_rows_for_smoke_runs(tmp_path, monkeypatch) -> None: + input_path = tmp_path / "train.jsonl" + output_path = tmp_path / "prepared.jsonl" + input_path.write_text( + "\n".join( + json.dumps( + { + "idx": index, + "dataset": "sciknoweval", + "kind": "mcq", + "answer": "A", + "prompt": f"Question {index}", + } + ) + for index in range(3) + ) + + "\n", + encoding="utf-8", + ) + monkeypatch.setattr( + sys, + "argv", + [ + "prepare_data", + "--dataset", + "sciknoweval", + "--input", + str(input_path), + "--domain", + "chemistry", + "--source-split", + "train", + "--max-rows", + "1", + "--output", + str(output_path), + ], + ) + + main() + + assert len(output_path.read_text(encoding="utf-8").splitlines()) == 1 + + +def test_prepare_data_main_limits_rows_after_filtering(tmp_path, monkeypatch) -> None: + input_path = tmp_path / "train.jsonl" + output_path = tmp_path / "prepared.jsonl" + input_path.write_text( + "\n".join( + json.dumps( + { + "idx": index, + "question": f"Question {index}", + "answerKey": "A", + "choices": {"label": ["A"], "text": ["answer"]}, + "domain": "Physics", + "details": {"level": level}, + } + ) + for index, level in enumerate(("L2", "L3", "L3")) + ) + + "\n", + encoding="utf-8", + ) + monkeypatch.setattr( + sys, + "argv", + [ + "prepare_data", + "--dataset", + "sciknoweval", + "--input", + str(input_path), + "--source-split", + "test", + "--max-rows", + "1", + "--output", + str(output_path), + ], + ) + + main() + + rows = [json.loads(line) for line in output_path.read_text(encoding="utf-8").splitlines()] + assert len(rows) == 1 + assert rows[0]["prompt"].startswith("Question 1") + + +def test_normalize_sciknoweval_uses_answer_when_answer_key_is_empty() -> None: + from examples.on_policy_distillation.sdpo.prepare_data import normalize_rows + + normalized = normalize_rows( + "sciknoweval", + [ + { + "question": "Fill it.", + "answerKey": "", + "answer": "H2O", + "domain": "Chemistry", + "details": {"level": "L3"}, + } + ], + source_split="test", + ) + + assert normalized[0]["label"] == "H2O" + assert normalized[0]["metadata"]["answer_key"] == "H2O" diff --git a/tests/utils/data/test_data_utils.py b/tests/utils/data/test_data_utils.py index c9ed438de..74bc0167a 100644 --- a/tests/utils/data/test_data_utils.py +++ b/tests/utils/data/test_data_utils.py @@ -4,7 +4,7 @@ import pytest -from relax.utils.data.data_utils import build_messages, collect_message_multimodal_data +from relax.utils.data.data_utils import build_messages, collect_message_multimodal_data, process_raw_sample @pytest.mark.parametrize( @@ -75,3 +75,38 @@ def test_build_messages_uses_top_level_media_without_mutating_input(): assert first_messages == second_messages assert first_messages[0]["content"][0] == {"type": "image", "image": "/data/image.png"} assert collect_message_multimodal_data(first_messages)["image"] == ["/data/image.png"] + + +class _Tokenizer: + def apply_chat_template(self, messages, **kwargs): + return f"rendered:{messages[-1]['content']}" + + +def test_process_raw_sample_surfaces_teacher_prompt_as_metadata(): + sample = process_raw_sample( + {"text": "question", "teacher_text": "privileged question", "label": "answer"}, + _Tokenizer(), + processor=None, + prompt_key="text", + label_key="label", + teacher_prompt_key="teacher_text", + apply_chat_template=True, + ) + + assert sample.teacher_prompt is None + assert sample.metadata["opd_teacher_prompt"] == "rendered:privileged question" + + +def test_process_raw_sample_without_teacher_key_leaves_no_privilege_metadata(): + sample = process_raw_sample( + {"text": "question", "label": "answer"}, + _Tokenizer(), + processor=None, + prompt_key="text", + label_key="label", + teacher_prompt_key=None, + apply_chat_template=True, + ) + + assert sample.teacher_prompt is None + assert "opd_teacher_prompt" not in sample.metadata diff --git a/tests/utils/opd/test_feedback.py b/tests/utils/opd/test_feedback.py new file mode 100644 index 000000000..0c59e75a3 --- /dev/null +++ b/tests/utils/opd/test_feedback.py @@ -0,0 +1,487 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""Coverage for the OPD feedback strategy hierarchy (OPD / OPSD / SDPO).""" + +from __future__ import annotations + +from types import SimpleNamespace + +import numpy as np +import pytest + +from relax.utils.types import Sample + + +def _sample(group: int | None, index: int, response: str, reward: object, **metadata: object) -> Sample: + return Sample( + group_index=group, + index=index, + prompt=f"question-{group}", + response=response, + response_length=len(response), + reward=reward, + metadata=dict(metadata), + ) + + +_SDPO_FEEDBACK_CLASSES = ["GoldenAnswerSDPOFeedback"] + + +def _sdpo_feedback_class(name: str) -> type: + import relax.utils.opd.sdpo.feedback as sdpo_feedback_module + + return getattr(sdpo_feedback_module, name) + + +# --- record_sample_feedback ------------------------------------------------- + + +def test_record_appends_text_only_to_originating_sample() -> None: + from relax.utils.opd.feedback import EnvironmentFeedback + + sample = _sample(1, 0, "bad", 0.0) + other = _sample(1, 1, "good", 1.0) + EnvironmentFeedback.record(sample, "failed test case") + EnvironmentFeedback.record(sample, "retry with a shorter answer") + + assert sample.metadata["env_feedback"] == ["failed test case", "retry with a shorter answer"] + assert "env_feedback" not in other.metadata + + +def test_record_sample_feedback_defaults_to_env_feedback_recording() -> None: + from relax.utils.opd.feedback import OPDFeedback + from relax.utils.opd.opsd.feedback import OPSDFeedback + from relax.utils.opd.sdpo.feedback import GoldenAnswerSDPOFeedback + + opd_sample = _sample(1, 0, "answer", {"score": 0.0, "feedback": "opd said no"}) + opsd_sample = _sample(1, 1, "answer", {"score": 0.0, "feedback": "opsd said no"}) + sdpo_sample = _sample(1, 2, "answer", {"score": 0.0, "feedback": "fix the second step"}) + + OPDFeedback().record_sample_feedback(opd_sample, opd_sample.reward) + OPSDFeedback(teacher_prompt_key="teacher_prompt").record_sample_feedback(opsd_sample, opsd_sample.reward) + GoldenAnswerSDPOFeedback().record_sample_feedback(sdpo_sample, sdpo_sample.reward) + + assert opd_sample.metadata["env_feedback"] == ["opd said no"] + assert opsd_sample.metadata["env_feedback"] == ["opsd said no"] + assert sdpo_sample.metadata["env_feedback"] == ["fix the second step"] + + +# --- teacher prompt construction -------------------------------------------- + + +@pytest.mark.parametrize("feedback_class_name", _SDPO_FEEDBACK_CLASSES) +def test_sdpo_peer_solution_shared_only_inside_group_and_normal_error_dropped(feedback_class_name: str) -> None: + target = _sample(7, 0, "wrong", {"score": 0.0}, env_feedback=["fix arithmetic"]) + success = _sample(7, 1, "correct solution", {"score": 1.0}, env_feedback=["success details"]) + unrelated = _sample(8, 2, "other", {"score": 0.0}, env_feedback=["unrelated feedback"]) + + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts( + [target, success, unrelated], [target.reward, success.reward, unrelated.reward] + ) + + assert target.teacher_prompt is not None + target_text = ( + target.teacher_prompt[-1]["content"] if isinstance(target.teacher_prompt, list) else target.teacher_prompt + ) + assert "correct solution" in target_text + assert "fix arithmetic" not in target_text + assert "success details" not in target_text + assert unrelated.teacher_prompt == unrelated.prompt + assert unrelated.opd_sample_mask is False + assert "unrelated feedback" not in str(unrelated.teacher_prompt) + + +@pytest.mark.parametrize("feedback_class_name", _SDPO_FEEDBACK_CLASSES) +def test_sdpo_prefers_peer_solution_and_falls_back_to_self_success(feedback_class_name: str) -> None: + self_success = _sample(9, 0, "self solution", {"score": 1.0}) + peer_success = _sample(9, 1, "peer solution", {"score": 1.0}) + + singleton_success = _sample(10, 0, "only solution", {"score": 1.0}) + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts( + [self_success, peer_success, singleton_success], + [self_success.reward, peer_success.reward, singleton_success.reward], + ) + + assert "peer solution" in str(self_success.teacher_prompt) + assert "self solution" not in str(self_success.teacher_prompt) + assert "self solution" in str(peer_success.teacher_prompt) + assert "peer solution" not in str(peer_success.teacher_prompt) + assert "only solution" in str(singleton_success.teacher_prompt) + assert singleton_success.opd_sample_mask is True + + +def test_sdpo_falls_back_to_original_prompt_without_solution_or_feedback() -> None: + from relax.utils.opd.sdpo.feedback import GoldenAnswerSDPOFeedback + + sample = _sample(3, 0, "an answer", {"score": 0.0}) + GoldenAnswerSDPOFeedback().prepare_teacher_prompts([sample], [sample.reward]) + + assert sample.teacher_prompt == sample.prompt + assert sample.opd_sample_mask is False + + +def test_sdpo_fallback_copies_original_message_prompt() -> None: + from relax.utils.opd.sdpo.feedback import GoldenAnswerSDPOFeedback + + prompt = [{"role": "user", "content": "question"}] + sample = Sample( + group_index=3, + prompt=prompt, + response="answer", + response_length=1, + reward={"score": 0.0}, + ) + + GoldenAnswerSDPOFeedback().prepare_teacher_prompts([sample], [sample.reward]) + + assert sample.teacher_prompt == prompt + assert sample.teacher_prompt is not prompt + assert sample.opd_sample_mask is False + + +@pytest.mark.parametrize("feedback_class_name", _SDPO_FEEDBACK_CLASSES) +def test_sdpo_format_error_without_peer_injects_feedback(feedback_class_name: str) -> None: + sample = _sample(12, 0, "wrong", {"score": 0.0, "format_error": 1, "feedback": "wrong format message"}) + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts([sample], [sample.reward]) + + assert "wrong format message" in str(sample.teacher_prompt) + assert sample.opd_sample_mask is True + + +@pytest.mark.parametrize("feedback_class_name", _SDPO_FEEDBACK_CLASSES) +def test_sdpo_truncation_without_peer_injects_feedback(feedback_class_name: str) -> None: + sample = _sample(13, 0, "truncated", {"score": 0.0, "feedback": "truncation message"}) + sample.status = Sample.Status.TRUNCATED + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts([sample], [sample.reward]) + + assert "truncation message" in str(sample.teacher_prompt) + assert sample.opd_sample_mask is True + + +@pytest.mark.parametrize("feedback_class_name", _SDPO_FEEDBACK_CLASSES) +def test_sdpo_normal_wrong_without_peer_drops_feedback_and_skips_distillation(feedback_class_name: str) -> None: + sample = _sample(14, 0, "wrong", {"score": 0.0, "feedback": "generic wrong-answer feedback"}) + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts([sample], [sample.reward]) + + assert sample.teacher_prompt == sample.prompt + assert sample.opd_sample_mask is False + + +@pytest.mark.parametrize("feedback_class_name", _SDPO_FEEDBACK_CLASSES) +def test_sdpo_uses_same_group_successful_rollout(feedback_class_name: str) -> None: + failed = _sample(4, 0, "wrong", {"score": 0.0}) + solved = _sample(4, 1, "worked solution", {"score": 1.0}) + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts( + [failed, solved], [failed.reward, solved.reward] + ) + + assert "worked solution" in str(failed.teacher_prompt) + assert failed.opd_sample_mask is True + + +@pytest.mark.parametrize("feedback_class_name", ["GoldenAnswerSDPOFeedback"]) +def test_tool_use_uid_when_group_index_is_missing(feedback_class_name: str) -> None: + failed = _sample(None, 0, "failed attempt", {"score": 0.0}, uid="uid-a") + solved = _sample(None, 1, "same uid solution", {"score": 1.0}, uid="uid-a") + unrelated = _sample(None, 2, "other uid solution", {"score": 1.0}, uid="uid-b") + + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts( + [failed, solved, unrelated], [failed.reward, solved.reward, unrelated.reward] + ) + + assert "same uid solution" in str(failed.teacher_prompt) + assert "other uid solution" not in str(failed.teacher_prompt) + assert failed.opd_sample_mask is True + + +def test_sdpo_success_reward_threshold_configures_success_boundary() -> None: + from relax.utils.opd.sdpo.feedback import GoldenAnswerSDPOFeedback + + low = _sample(5, 0, "partial solution", {"score": 0.9}) + strict = GoldenAnswerSDPOFeedback(success_reward_threshold=1.0) + strict.prepare_teacher_prompts([low], [low.reward]) + assert low.teacher_prompt == low.prompt + assert low.opd_sample_mask is False + + lenient = GoldenAnswerSDPOFeedback(success_reward_threshold=0.8) + lenient.prepare_teacher_prompts([low], [low.reward]) + assert "partial solution" in str(low.teacher_prompt) + assert low.opd_sample_mask is True + + +def test_opsd_assigns_dataset_privilege_while_opd_stays_empty() -> None: + from relax.utils.opd.feedback import OPDFeedback + from relax.utils.opd.opsd.feedback import OPSDFeedback + + opd_sample = _sample(6, 0, "answer", 1.0) + privileged = _sample(6, 1, "answer", 1.0) + privileged.metadata["opd_teacher_prompt"] = "dataset teacher prompt" + unprivileged = _sample(6, 2, "answer", 1.0) + + OPDFeedback().prepare_teacher_prompts([opd_sample], [opd_sample.reward]) + OPSDFeedback(teacher_prompt_key="teacher_prompt").prepare_teacher_prompts( + [privileged, unprivileged], [privileged.reward, unprivileged.reward] + ) + + assert opd_sample.teacher_prompt is None + assert opd_sample.opd_sample_mask is None + assert privileged.teacher_prompt == "dataset teacher prompt" + assert privileged.opd_sample_mask is None + assert unprivileged.teacher_prompt is None + assert unprivileged.opd_sample_mask is None + + +@pytest.mark.parametrize("feedback_class_name", ["GoldenAnswerSDPOFeedback"]) +def test_tool_use_feedback_share_peer_solutions(feedback_class_name: str) -> None: + failed = _sample(11, 0, "failed attempt", {"score": 0.0, "feedback": "fix the tool call"}) + solved = _sample(11, 1, "successful attempt", {"score": 1.0}) + _sdpo_feedback_class(feedback_class_name)().prepare_teacher_prompts( + [failed, solved], [failed.reward, solved.reward] + ) + + teacher_text = str(failed.teacher_prompt) + assert "successful attempt" in teacher_text + assert "fix the tool call" not in teacher_text + assert failed.opd_sample_mask is True + + +# --- strategy hook contract --------------------------------------------------- + + +def test_load_feedback_class_defaults_to_opd() -> None: + from relax.utils.opd.feedback import OPDFeedback, load_feedback_class + + assert load_feedback_class(None) is OPDFeedback + assert load_feedback_class("") is OPDFeedback + + +def test_load_feedback_class_rejects_non_feedback_path() -> None: + from relax.utils.opd.feedback import load_feedback_class + + with pytest.raises(TypeError, match="EnvironmentFeedback subclass"): + load_feedback_class("relax.utils.opd.opd_utils.OPD_SAMPLE_MASK") + + with pytest.raises(ValueError, match="Invalid feedback class path"): + load_feedback_class("not-a-dotted-path") + + +def test_load_feedback_defaults_to_opd_with_empty_kwargs() -> None: + from relax.utils.opd.feedback import OPDFeedback, load_feedback + + feedback = load_feedback(None, None) + assert isinstance(feedback, OPDFeedback) + assert feedback.teacher_prompt_key is None + assert feedback.success_reward_threshold == 1.0 + + +def test_load_feedback_binds_kwargs_to_selected_class() -> None: + from relax.utils.opd.feedback import load_feedback + from relax.utils.opd.opsd.feedback import OPSDFeedback + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + sdpo = load_feedback( + "relax.utils.opd.sdpo.feedback.SDPOFeedback", + {"success_reward_threshold": 0.8}, + ) + assert isinstance(sdpo, SDPOFeedback) + assert sdpo.success_reward_threshold == 0.8 + + opsd = load_feedback( + "relax.utils.opd.opsd.feedback.OPSDFeedback", + {"teacher_prompt_key": "teacher_prompt"}, + ) + assert isinstance(opsd, OPSDFeedback) + assert opsd.teacher_prompt_key == "teacher_prompt" + + +def test_load_feedback_rejects_unknown_kwargs_and_missing_opsd_key() -> None: + from relax.utils.opd.feedback import load_feedback + + with pytest.raises(TypeError, match="Invalid --opd-feedback-kwargs for SDPOFeedback"): + load_feedback("relax.utils.opd.sdpo.feedback.SDPOFeedback", {"unknown_param": 1}) + + with pytest.raises(ValueError, match="requires"): + load_feedback("relax.utils.opd.opsd.feedback.OPSDFeedback", {}) + + +def test_code_sdpo_feedback_is_a_placeholder() -> None: + from relax.utils.opd.sdpo.feedback import CodeSDPOFeedback + + feedback = CodeSDPOFeedback() + with pytest.raises(NotImplementedError): + feedback.record_sample_feedback(Sample(), 0.0) + with pytest.raises(NotImplementedError): + feedback.prepare_teacher_prompts([Sample()], [0.0]) + + +def test_opd_feedback_hooks_are_no_ops() -> None: + from relax.utils.opd.feedback import OPDFeedback + + feedback = OPDFeedback() + assert feedback.extra_transfer_schema() == [] + feedback.produce_extra_transfer([Sample()], {}) + feedback.check_student_topk_ids(Sample(), 8) + feedback.check_transfer_channels(Sample(), {}, 8) + assert OPDFeedback.validate_launch_args(SimpleNamespace()) is None + + +def test_opd_create_opsd_worker_follows_dataset_keys() -> None: + from relax.utils.opd.feedback import OPDFeedback + + args = SimpleNamespace(opd_teacher_image_key=None) + assert OPDFeedback.create_opsd_worker(args).is_opsd is False + + args.opd_teacher_image_key = "images" + assert OPDFeedback.create_opsd_worker(args).is_opsd is True + + +def test_sdpo_extra_transfer_schema_and_mask_column() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + feedback = SDPOFeedback() + assert feedback.extra_transfer_schema() == ["opd_sample_mask"] + + samples = [Sample(index=0, opd_sample_mask=True), Sample(index=1, opd_sample_mask=False)] + train_data: dict = {} + feedback.produce_extra_transfer(samples, train_data) + + assert train_data["opd_sample_mask"] == [True, False] + + +def test_sdpo_create_opsd_worker_is_always_active() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + assert SDPOFeedback.create_opsd_worker(SimpleNamespace()).is_opsd is True + + +def _sdpo_launch_args(**overrides: object) -> SimpleNamespace: + args = SimpleNamespace( + opd_type="sglang", + pipeline_model_parallel_size=1, + enable_mtp_training=False, + opd_token_selection="student_topk", + opd_kl_type="jsd", + opd_norm_mode="tail", + calculate_per_token_loss=True, + multimodal_keys=None, + opd_teacher_image_key=None, + opd_teacher_video_key=None, + opd_teacher_audio_key=None, + group_rm=True, + opd_kl_coef=0.0, + opd_loss_coef=1.0, + ) + for key, value in overrides.items(): + setattr(args, key, value) + return args + + +def test_sdpo_validate_launch_args_accepts_valid_configuration() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + assert SDPOFeedback.validate_launch_args(_sdpo_launch_args()) is None + + +@pytest.mark.parametrize( + ("overrides", "message"), + [ + ({"group_rm": False}, "group-rm"), + ({"opd_token_selection": "teacher_topk"}, "student_topk"), + ({"opd_kl_coef": 1.0}, "opd-loss-coef"), + ({"opd_type": "megatron"}, "opd-type=sglang"), + ({"pipeline_model_parallel_size": 2}, "pipeline parallelism"), + ({"calculate_per_token_loss": False}, "calculate-per-token-loss"), + ], +) +def test_sdpo_validate_launch_args_rejects_invalid_configuration(overrides: dict, message: str) -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + with pytest.raises(ValueError, match=message): + SDPOFeedback.validate_launch_args(_sdpo_launch_args(**overrides)) + + +def test_sdpo_check_student_topk_ids_delegates_to_validation() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + feedback = SDPOFeedback() + valid = Sample(index=0, response_length=2, student_topk_token_ids=np.array([[1, 2], [3, 4]], dtype=np.int32)) + feedback.check_student_topk_ids(valid, 2) + + bad_shape = Sample(index=0, response_length=2, student_topk_token_ids=np.array([[1, 2]], dtype=np.int32)) + with pytest.raises(ValueError, match="student Top-K"): + feedback.check_student_topk_ids(bad_shape, 2) + + with pytest.raises(ValueError, match="student Top-K"): + feedback.check_student_topk_ids(Sample(index=0, response_length=2), 2) + + +def test_sdpo_check_transfer_channels_delegates_to_validation() -> None: + from relax.utils.opd.opd_main_worker import TopkWorker + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + feedback = SDPOFeedback() + feedback.check_transfer_channels(Sample(index=0, response_length=0), {}, 2) + + non_empty = Sample(index=0, response_length=1) + with pytest.raises(ValueError, match="complete teacher Top-K payload"): + feedback.check_transfer_channels(non_empty, {}, 2) + + complete = { + TopkWorker.TRANSFER_TOKEN_IDS: np.array([[1, 2]], dtype=np.int32), + TopkWorker.TRANSFER_TEACHER_LOG_PROBS: np.array([[-0.1, -0.2]], dtype=np.float32), + } + feedback.check_transfer_channels(non_empty, complete, 2) + + +def test_sdpo_prepare_teacher_prompts_rejects_multimodal() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + sample = Sample( + prompt="question", + response="answer", + response_length=1, + multimodal_inputs={"images": [b"image"]}, + ) + + with pytest.raises(ValueError, match="SDPO only supports text inputs"): + SDPOFeedback().prepare_teacher_prompts([sample], [{"score": 0.0}]) + + +def test_sdpo_prepare_teacher_prompts_requires_teacher_prompt_for_nonempty_response() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + empty_prompt = Sample(prompt="", response="answer", response_length=6, reward={"score": 0.0}) + with pytest.raises(ValueError, match="SDPO requires a teacher prompt"): + SDPOFeedback().prepare_teacher_prompts([empty_prompt], [empty_prompt.reward]) + + empty_response = Sample(prompt="", response="", response_length=0, reward={"score": 0.0}) + SDPOFeedback().prepare_teacher_prompts([empty_response], [empty_response.reward]) + + +def test_sdpo_prepare_teacher_prompts_clears_stale_teacher_payload() -> None: + from relax.utils.opd.sdpo.feedback import SDPOFeedback + + sample = Sample( + group_index=1, + prompt="question", + response="answer", + response_length=1, + reward={"score": 0.0}, + student_topk_token_ids=np.array([[1, 2]], dtype=np.int32), + student_topk_log_probs=np.array([[-0.1, -0.2]], dtype=np.float32), + teacher_log_probs=[-0.3], + teacher_topk_token_ids=np.array([[3, 4]], dtype=np.int32), + opd_topk_token_ids=np.array([[3, 4]], dtype=np.int32), + teacher_tokens=[9, 10], + teacher_prompt_length=1, + ) + + SDPOFeedback().prepare_teacher_prompts([sample], [sample.reward]) + + assert sample.teacher_log_probs is None + assert sample.teacher_topk_token_ids is None + assert sample.opd_topk_token_ids is None + assert sample.teacher_tokens is None + assert sample.teacher_prompt_length is None + np.testing.assert_array_equal(sample.student_topk_token_ids, np.array([[1, 2]], dtype=np.int32)) + np.testing.assert_array_equal(sample.student_topk_log_probs, np.array([[-0.1, -0.2]], dtype=np.float32)) diff --git a/tests/utils/opd/test_opd_topk_log_probs.py b/tests/utils/opd/test_opd_topk_log_probs.py new file mode 100644 index 000000000..555ff8dc8 --- /dev/null +++ b/tests/utils/opd/test_opd_topk_log_probs.py @@ -0,0 +1,29 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +"""Regression tests for global Top-K token ids across TP shards.""" + +import pytest +import torch + +from relax.utils.opd.opd_utils import compute_log_probs_on_topk_token_ids # noqa: E402 + + +def test_topk_log_probs_keep_global_zero_and_sentinel_ids() -> None: + logits = torch.tensor([[0.0, 1.0, 2.0, 3.0]], requires_grad=True) + token_ids = torch.tensor([[0, 3, -1]], dtype=torch.long) + + actual = compute_log_probs_on_topk_token_ids(logits, token_ids, process_group=None) + expected = torch.cat([logits[:, :1], logits[:, 3:4], torch.full((1, 1), float("-inf"))], dim=-1) - torch.logsumexp( + logits, dim=-1, keepdim=True + ) + + assert torch.allclose(actual[:, :2], expected[:, :2]) + assert torch.isneginf(actual[:, 2]).all() + + +def test_topk_log_probs_reject_ids_outside_global_vocabulary() -> None: + logits = torch.zeros((1, 4)) + token_ids = torch.tensor([[4]], dtype=torch.long) + + with pytest.raises(ValueError, match="global vocabulary size"): + compute_log_probs_on_topk_token_ids(logits, token_ids, process_group=None) diff --git a/tests/utils/test_arguments_opd_teacher_colocate.py b/tests/utils/test_arguments_opd_teacher_colocate.py index b87721536..b98e984c2 100644 --- a/tests/utils/test_arguments_opd_teacher_colocate.py +++ b/tests/utils/test_arguments_opd_teacher_colocate.py @@ -97,16 +97,16 @@ def _opd_args() -> SimpleNamespace: use_opd=True, opd_kl_coef=1.0, opd_loss_coef=0.0, - opd_teacher_prompt_key=None, + calculate_per_token_loss=True, + opd_feedback_class="relax.utils.opd.feedback.OPDFeedback", + opd_feedback_kwargs=None, opd_teacher_image_key=None, opd_teacher_video_key=None, opd_teacher_audio_key=None, multimodal_keys=None, opd_per_token_clip=None, opd_is_clip=None, - opd_mask_on_success=False, opd_only_reward=False, - opd_log_prob_dump_dir=None, opd_type="sglang", opd_teacher_load=None, teacher_hf_checkpoint="/teacher", @@ -170,6 +170,7 @@ def _opd_args() -> SimpleNamespace: rollout_max_context_len=None, rollout_max_prompt_len=None, qkv_format="sbhd", + allgather_cp=False, train_backend="megatron", only_train_params_name_list=None, freeze_params_name_list=None, @@ -194,6 +195,46 @@ def test_opd_sampled_token_loss_is_accepted(arguments_module): arguments_module.slime_validate_args(args) +def test_opd_jsd_help_describes_existing_alpha(arguments_module): + arguments_module.RouterArgs = SimpleNamespace(add_cli_args=lambda parser, **_kwargs: parser) + parser = argparse.ArgumentParser() + arguments_module.get_slime_extra_args_provider()(parser) + + help_text = next(action.help for action in parser._actions if action.dest == "opd_jsd_alpha") + + assert "Mixture coefficient" in help_text + + +def test_topk_opd_rejects_allgather_context_parallel(arguments_module): + args = _opd_args() + args.opd_token_selection = "student_topk" + args.opd_log_prob_top_k = 2 + args.allgather_cp = True + + with pytest.raises(ValueError, match="not compatible with --allgather-cp"): + arguments_module.slime_validate_args(args) + + +def test_topk_opd_rejects_context_parallelism(arguments_module): + args = _opd_args() + args.opd_token_selection = "student_topk" + args.opd_log_prob_top_k = 2 + args.context_parallel_size = 2 + + with pytest.raises(ValueError, match="not compatible with context parallelism"): + arguments_module.slime_validate_args(args) + + +def test_topk_opd_rejects_dynamic_context_parallel(arguments_module): + args = _opd_args() + args.opd_token_selection = "student_topk" + args.opd_log_prob_top_k = 2 + args.dynamic_context_parallel = True + + with pytest.raises(ValueError, match="not compatible with context parallelism"): + arguments_module.slime_validate_args(args) + + def test_managed_opd_teacher_colocate_preserves_rollout_resource_split(arguments_module): args = _opd_args() args.colocate = True diff --git a/tests/utils/test_dict_to_tensordict_ragged.py b/tests/utils/test_dict_to_tensordict_ragged.py new file mode 100644 index 000000000..1069952a5 --- /dev/null +++ b/tests/utils/test_dict_to_tensordict_ragged.py @@ -0,0 +1,55 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +import pytest +import torch + +from relax.utils.utils import dict_to_tensordict + + +def test_dict_to_tensordict_preserves_mixed_empty_and_nonempty_rows() -> None: + result = dict_to_tensordict({"opd_topk_teacher_log_probs": [[], [-1.0, -2.0]]}, batch_size=2) + + rows = result["opd_topk_teacher_log_probs"] + assert rows.layout == torch.jagged + assert rows[0].numel() == 0 + assert torch.equal(rows[1], torch.tensor([-1.0, -2.0])) + + +def test_dict_to_tensordict_preserves_nonempty_row_before_empty_row() -> None: + result = dict_to_tensordict({"opd_topk_teacher_log_probs": [[-1.0, -2.0], []]}, batch_size=2) + + rows = result["opd_topk_teacher_log_probs"] + assert rows[0].tolist() == [-1.0, -2.0] + assert rows[1].numel() == 0 + assert rows[0].dtype == rows[1].dtype + + +def test_dict_to_tensordict_preserves_all_empty_ragged_rows() -> None: + result = dict_to_tensordict({"opd_topk_token_ids": [[], []]}, batch_size=2) + + rows = result["opd_topk_token_ids"] + assert rows.layout == torch.jagged + assert rows[0].numel() == 0 + assert rows[1].numel() == 0 + + +@pytest.mark.parametrize( + ("field", "values"), + [ + ("opd_topk_token_ids", [11, 12]), + ("opd_topk_ksz", [2]), + ], +) +@pytest.mark.parametrize("empty_first", [True, False]) +def test_dict_to_tensordict_preserves_mixed_integer_rows(field, values, empty_first) -> None: + rows = [[], values] if empty_first else [values, []] + + result = dict_to_tensordict({field: rows}, batch_size=2) + converted = result[field] + nonempty_index = 1 if empty_first else 0 + empty_index = 0 if empty_first else 1 + + assert converted.layout == torch.jagged + assert converted[empty_index].numel() == 0 + assert converted[nonempty_index].tolist() == values + assert converted[empty_index].dtype == converted[nonempty_index].dtype