Skip to content

Commit 7dfe73c

Browse files
Merge pull request #4362 from AI-Hypercomputer:igorts/dpo-notebook-with-eval
PiperOrigin-RevId: 948420388
2 parents da18954 + f5b5888 commit 7dfe73c

5 files changed

Lines changed: 4693 additions & 2 deletions

File tree

.github/workflows/run_jupyter_notebooks.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -111,7 +111,7 @@ jobs:
111111
112112
for notebook in "$MAXTEXT_NOTEBOOKS_ROOT"/*.ipynb; do
113113
filename=$(basename "$notebook")
114-
if [[ "$filename" == "sft_llama3_demo_gpu.ipynb" || "$filename" == "maxtext_with_gepa.ipynb" || "$filename" == "demo_decoding.ipynb" ]]; then
114+
if [[ "$filename" == "sft_llama3_demo_gpu.ipynb" || "$filename" == "maxtext_with_gepa.ipynb" || "$filename" == "demo_decoding.ipynb" || "$filename" == "dpo_qwen3_demo.ipynb" ]]; then
115115
echo "Skipping $filename"
116116
continue
117117
fi

docs/guides/run_python_notebook.md

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -186,6 +186,10 @@ jupyter lab --ip=0.0.0.0 --port=8888 --no-browser --allow-root
186186
- **`sft_qwen3_demo.ipynb`** → Qwen3-0.6B SFT training and evaluation on [OpenAI's GSM8K dataset](https://huggingface.co/datasets/openai/gsm8k). This notebook is friendly for beginners and runs successfully on Google Colab's free-tier v5e-1 TPU runtime.
187187
- **`sft_llama3_demo_tpu.ipynb`** → Llama3.1-8B SFT training on [Hugging Face ultrachat_200k dataset](https://huggingface.co/datasets/HuggingFaceH4/ultrachat_200k). We recommend running this on a v5p-8 TPU VM using [Method 2](#method-2-visual-studio-code-with-tpu-recommended) or [Method 3](#method-3-local-jupyter-lab-with-tpu-recommended).
188188

189+
### Direct Preference Optimization (DPO)
190+
191+
- **`dpo_qwen3_demo.ipynb`** → Qwen3-0.6B DPO training and evaluation on [Argilla's distilabel-intel-orca-dpo-pairs dataset](https://huggingface.co/datasets/argilla/distilabel-intel-orca-dpo-pairs). This notebook is fully optimized to run evaluations across all 4 TPUs using Tensor Parallelism (TP=4) on a TPU v5p-8 VM.
192+
189193
### Reinforcement Learning (GRPO/GSPO) Training
190194

191195
- **`rl_llama3_demo.ipynb`** → GRPO/GSPO training on [OpenAI's GSM8K dataset](https://huggingface.co/datasets/openai/gsm8k). We recommend running this on a v5p-8 TPU VM using [Method 2](#method-2-visual-studio-code-with-tpu-recommended) or [Method 3](#method-3-local-jupyter-lab-with-tpu-recommended).

src/maxtext/eval/runner/harness_runner.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -74,7 +74,23 @@ def _map_results(raw_results: dict, tasks: list[str]) -> dict:
7474
if acc_norm is not None:
7575
scores[f"{task}_accuracy_norm"] = round(float(acc_norm) * 100, 2)
7676

77-
if acc is None and task_r:
77+
# Extract ifeval and other task-specific custom accuracy keys
78+
has_custom_keys = False
79+
for suffix in (
80+
"prompt_level_strict_acc,none",
81+
"inst_level_strict_acc,none",
82+
"prompt_level_loose_acc,none",
83+
"inst_level_loose_acc,none",
84+
"prompt_level_strict_acc",
85+
"inst_level_strict_acc",
86+
"prompt_level_loose_acc",
87+
"inst_level_loose_acc",
88+
):
89+
if task_r.get(suffix) is not None:
90+
scores[f"{task}_{suffix.replace(',none', '')}"] = round(float(task_r[suffix]) * 100, 2)
91+
has_custom_keys = True
92+
93+
if acc is None and not has_custom_keys and task_r:
7894
logger.warning(
7995
"No known accuracy keys found for task '%s'. Available: %s",
8096
task,
@@ -167,6 +183,8 @@ def run_harness(cfg: dict, hf_token: str | None = None) -> dict:
167183
"limit": num_samples,
168184
"log_samples": False,
169185
}
186+
if cfg.get("max_num_seqs") is not None:
187+
simple_eval_kwargs["batch_size"] = cfg["max_num_seqs"]
170188
if apply_chat_template:
171189
simple_eval_kwargs["apply_chat_template"] = True
172190
if fewshot_as_multiturn:

0 commit comments

Comments
 (0)