diff --git a/examples/README.md b/examples/README.md index 1daea83c0..f8ea2cb74 100644 --- a/examples/README.md +++ b/examples/README.md @@ -11,7 +11,7 @@ These examples provide concrete examples to leverage vime in your own RL workflo - **[low_precision](./low_precision)**: Examples of FP8 training and inference for improved throughput and stability. - **[mem_agent](./mem_agent)**: MemAgent long-context RL — chunk-wise memory update, HotpotQA GRPO training, and RULER-HQA evaluation. - **[multi_agent](./multi_agent)**: Example of running multi-agent RL with `vime`. -- **[on_policy_distillation](./on_policy_distillation)**: Example implementation for on-policy distillation, extending the reinforcement learning pipeline to support teacher–student distillation directly within on-policy training. +- **[on_policy_distillation](./on_policy_distillation)**: On-policy distillation (OPD) with an external vLLM teacher or a Megatron-loaded teacher. - **[delta_weight_sync](./delta_weight_sync)**: Non-colocated weight sync that ships only the changed bytes over a shared filesystem (training/inference disaggregation), reloading via the vanilla `update_weights_from_disk` path. - **[reproducibility](./reproducibility)**: Guides on achieving bitwise experiment reproduction using deterministic modes. - **[retool](./retool)**: Demonstrates the retool functionality for tool-enabled language model generation. diff --git a/examples/on_policy_distillation/README.md b/examples/on_policy_distillation/README.md new file mode 100644 index 000000000..9d85d3e69 --- /dev/null +++ b/examples/on_policy_distillation/README.md @@ -0,0 +1,168 @@ +# On-Policy Distillation Example + +This example shows how to run **on-policy distillation (OPD)** using vime. A +small student (Qwen3-8B) is aligned to imitate a larger teacher (Qwen3-32B) by +training only on the student's own rollouts and matching the teacher's +token-level log-probabilities. + +## Key Features + +- **OPD is orthogonal to advantage estimators**: OPD works as an additive KL + penalty on top of any advantage estimator (GRPO, PPO, REINFORCE++, etc.), not + as a separate estimator. +- **Two teacher modes**: + - **vllm**: Teacher runs on an external vLLM server; teacher log-probs are + obtained during rollout. + - **megatron**: Teacher is loaded directly into Megatron via + `--opd-teacher-load`; teacher log-probs are computed during the training + forward pass. +- **Student rollout always uses vLLM** (vime's default rollout backend). + +## Key Arguments + +| Argument | Description | +|----------|-------------| +| `--use-opd` | Enable on-policy distillation. Required flag to use OPD. | +| `--opd-type` | Type of OPD: `vllm` or `megatron`. Required when `--use-opd` is set. | +| `--opd-kl-coef` | OPD KL penalty coefficient (default: 1.0). | +| `--opd-teacher-load` | Path to teacher checkpoint. **Required** when `--opd-type=megatron`, **must not be set** when `--opd-type=vllm`. | +| `--opd-teacher-ckpt-step` | Optional checkpoint step for teacher model. | + +## Mode Comparison + +| Mode | Teacher Location | When to use | +|------|------------------|-------------| +| `vllm` | External vLLM server | Teacher has different architecture or is larger than GPU memory | +| `megatron` | Loaded into Megatron training | Teacher has same architecture as policy/ref model | + +## Components + +- `vime/rollout/on_policy_distillation.py` implements (for vLLM mode): + - `reward_func` calls the teacher server (via `args.rm_url`) with every sample + to obtain token-level logprobs. + - `post_process_rewards` trims the teacher logprobs to the generated response + span and writes the tensors back to each `Sample` to compute advantages. +- `run-qwen3-8B-opd.sh` launches a vLLM teacher server, then submits a Ray job + that runs `train.py`. +- `run-qwen3-8B-opd-megatron.sh` uses a Megatron-loaded teacher model (no + external server needed). + +## Running the example + +### Using vLLM Teacher (External Server) + +1. Download or prepare the required checkpoints and data. + +```bash +hf download Qwen/Qwen3-32B --local-dir /root/Qwen3-32B +hf download Qwen/Qwen3-8B --local-dir /root/Qwen3-8B +hf download --repo-type dataset zhuzilin/dapo-math-17k --local-dir /root/dapo-math-17k +``` + +2. Run the hf to mcore for student model conversion: + +```bash +cd /root/vime +source scripts/models/qwen3-8B.sh + +PYTHONPATH=/root/Megatron-LM python tools/convert_hf_to_torch_dist.py \ + ${MODEL_ARGS[@]} \ + --hf-checkpoint /root/Qwen3-8B \ + --save /root/Qwen3-8B_torch_dist +``` + +3. Run on-policy distillation: + +```bash +bash examples/on_policy_distillation/run-qwen3-8B-opd.sh +``` + +GPU layout: + +| GPUs | Role | +|------|------| +| 0–3 | Student Megatron train + student vLLM rollout (colocate) | +| 4–7 | Teacher vLLM (Qwen3-32B, TP=4) | + +### Using Megatron Teacher (No External Server) + +1. Prepare student checkpoint (same as above). + +2. **IMPORTANT**: Convert your teacher model to Megatron format (change the path + to your actual teacher): + +```bash +# This example uses the same model as both student and teacher (for demonstration only) +# In practice, use a different (stronger) model as the teacher! +cd /root/vime +source scripts/models/qwen3-8B.sh # Or your teacher model config + +PYTHONPATH=/root/Megatron-LM python tools/convert_hf_to_torch_dist.py \ + ${MODEL_ARGS[@]} \ + --hf-checkpoint /root/YourTeacherModel \ + --save /root/YourTeacherModel_torch_dist +``` + +3. Edit `run-qwen3-8B-opd-megatron.sh` to update paths: + - Change `--opd-teacher-load` to your teacher model path + - Adjust `--opd-kl-coef` based on your task + +4. Run: + +```bash +bash examples/on_policy_distillation/run-qwen3-8B-opd-megatron.sh +``` + +# Preliminary Results + +End-to-end run with `run-qwen3-8B-opd.sh` (dapo-math-17k train, GRPO + +`--opd-kl-coef 1.0`, ~220 rollouts / iter_0000219). Offline GSM8K greedy eval: + +| Model | GSM8K Accuracy | +|-------|----------------| +| Qwen3-8B (pre-OPD) | 79.7% (n=300) | +| Qwen3-8B (post-OPD) | **88.2%** (n=1319, **+8.5 pp**) | +| Qwen3-32B teacher | 87.0% (n=300) | + +Training health signal: `rollout/opd_reverse_kl` dropped from 0.145 → ~0.10 +(−38%). Pure OPD uses `raw_reward=0`; the learning signal is the OPD KL term. + +# FAQ + +1. **Why are there two OPD modes?** + - `vllm` mode: The teacher runs on an independent vLLM server. This is useful + when the teacher has a different architecture or is too large to load + together with the policy model. + - `megatron` mode: The teacher is loaded into Megatron using the same + parameter loading mechanism as the reference model. This requires the + teacher to have the same architecture as the policy model. + +2. **How do I use Megatron-based teacher instead of vLLM server?** + Replace your OPD arguments: + ```bash + # Instead of: + --use-opd --opd-type vllm --opd-kl-coef 1.0 + # Use: + --use-opd --opd-type megatron --opd-kl-coef 1.0 --opd-teacher-load /path/to/teacher_checkpoint + ``` + +3. **What happens if I set wrong arguments?** + The system will raise clear errors: + - `--use-opd` without `--opd-type`: Error asking you to specify type + - `--opd-type megatron` without `--opd-teacher-load`: Error asking for teacher checkpoint + - `--opd-type vllm` with `--opd-teacher-load`: Error indicating conflict + +4. **Why is `rollout/raw_reward` always 0?** + Pure OPD distillation does not use an external reward model. The learning + signal comes entirely from the OPD KL term applied to advantages. + +5. **Self-distillation: why is `opd_reverse_kl` near 0 at the start?** + Teacher and student start from the same weights, so reverse KL is ~0 until + the student updates. For a true distillation signal, use a stronger / + differently trained teacher (or `--opd-type vllm` with Qwen3-32B). + +# References + +1. https://thinkingmachines.ai/blog/on-policy-distillation/ +2. https://arxiv.org/abs/2306.13649 +3. https://arxiv.org/abs/2306.08543 diff --git a/examples/on_policy_distillation/run-qwen3-8B-opd-megatron.sh b/examples/on_policy_distillation/run-qwen3-8B-opd-megatron.sh new file mode 100644 index 000000000..22a0c1b2e --- /dev/null +++ b/examples/on_policy_distillation/run-qwen3-8B-opd-megatron.sh @@ -0,0 +1,162 @@ +#!/bin/bash + +# On-Policy Distillation with Megatron-based teacher model +# This example uses the original model as the teacher (self-distillation for demonstration) +# +# IMPORTANT: This is just an example configuration! +# In practice, you should: +# 1. Use a different (stronger) model as the teacher +# 2. Adjust --opd-kl-coef based on your task +# 3. Configure proper evaluation metrics + +set -ex + +export PYTHONUNBUFFERED=1 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +NVLINK_COUNT=$(nvidia-smi topo -m 2>/dev/null | grep -o 'NV[0-9][0-9]*' | wc -l) +if [ "$NVLINK_COUNT" -gt 0 ]; then + HAS_NVLINK=1 +else + HAS_NVLINK=0 +fi +echo "HAS_NVLINK: $HAS_NVLINK (detected $NVLINK_COUNT NVLink references)" + +source "/root/vime/scripts/models/qwen3-8B.sh" + +CKPT_ARGS=( + --hf-checkpoint /root/Qwen3-8B + --ref-load /root/Qwen3-8B_torch_dist + --load /root/Qwen3-8B_torch_dist + --save /root/Qwen3-8B_vime/ + --save-interval 10 + --megatron-to-hf-mode bridge +) + +ROLLOUT_ARGS=( + --prompt-data /root/dapo-math-17k/dapo-math-17k.jsonl + --input-key prompt + --apply-chat-template + --rollout-shuffle + --rm-type math + --num-rollout 300 + --rollout-batch-size 16 + --n-samples-per-prompt 4 + --rollout-max-response-len 16384 + --rollout-temperature 1 + + --global-batch-size 64 + --balance-data +) + +RM_ARGS=( + --rm-type math +) + +EVAL_ARGS=( + # --eval-interval 20 + # --eval-prompt-data aime /root/aime-2024/aime-2024.jsonl + # --n-samples-per-eval-prompt 16 + # --eval-max-response-len 16384 + # --eval-top-p 1 +) + +PERF_ARGS=( + --tensor-model-parallel-size 4 + --sequence-parallel + --pipeline-model-parallel-size 1 + --context-parallel-size 1 + --expert-model-parallel-size 1 + --expert-tensor-parallel-size 1 + + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 36 + + --use-dynamic-batch-size + --max-tokens-per-gpu 8192 +) + +GRPO_ARGS=( + --advantage-estimator grpo + # OPD Configuration + --use-opd + --opd-type megatron + --opd-kl-coef 1.0 + # CHANGE THIS to a stronger teacher checkpoint in practice + --opd-teacher-load /root/Qwen3-8B_torch_dist + --use-kl-loss + --kl-loss-coef 0.00 + --kl-loss-type low_var_kl + --entropy-coef 0.00 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.1 + --adam-beta1 0.9 + --adam-beta2 0.98 +) + +WANDB_ARGS=( + #--use-wandb + # --wandb-project vime-opd + # --wandb-group qwen3-8B-opd-megatron + # --wandb-key ${WANDB_KEY} +) + +VLLM_ARGS=( + --rollout-num-gpus-per-engine 1 + --vllm-gpu-memory-utilization 0.4 +) + +MISC_ARGS=( + --attention-dropout 0.0 + --hidden-dropout 0.0 + --accumulate-allreduce-grads-in-fp32 + --attention-softmax-in-fp32 + --attention-backend flash + --make-vocab-size-divisible-by 128 +) + +export MASTER_ADDR=${MASTER_ADDR:-"127.0.0.1"} +ray start --head --node-ip-address ${MASTER_ADDR} --num-gpus 8 --disable-usage-stats --dashboard-host=0.0.0.0 --dashboard-port=8265 + +RUNTIME_ENV_JSON="{ + \"env_vars\": { + \"PYTHONPATH\": \"/root/vime:/root/Megatron-LM/\", + \"CUDA_DEVICE_MAX_CONNECTIONS\": \"1\", + \"NCCL_NVLS_ENABLE\": \"${HAS_NVLINK}\", + \"PYTORCH_CUDA_ALLOC_CONF\": \"expandable_segments:True\" + } +}" + +ray job submit --address="http://127.0.0.1:8265" \ + --runtime-env-json="${RUNTIME_ENV_JSON}" \ + --working-dir /root/vime \ + -- python3 train.py \ + --train-backend megatron \ + --actor-num-nodes 1 \ + --actor-num-gpus-per-node 4 \ + --rollout-num-gpus 4 \ + --colocate \ + ${MODEL_ARGS[@]} \ + ${CKPT_ARGS[@]} \ + ${ROLLOUT_ARGS[@]} \ + ${OPTIMIZER_ARGS[@]} \ + ${GRPO_ARGS[@]} \ + ${WANDB_ARGS[@]} \ + ${PERF_ARGS[@]} \ + ${EVAL_ARGS[@]} \ + ${VLLM_ARGS[@]} \ + ${MISC_ARGS[@]} \ + ${RM_ARGS[@]} + +#### clear after training +sleep 3 +ray stop --force +pkill -9 ray +pkill -9 -f "train.py" || true +sleep 3 diff --git a/examples/on_policy_distillation/run-qwen3-8B-opd.sh b/examples/on_policy_distillation/run-qwen3-8B-opd.sh new file mode 100644 index 000000000..c02f6bc79 --- /dev/null +++ b/examples/on_policy_distillation/run-qwen3-8B-opd.sh @@ -0,0 +1,202 @@ +#!/bin/bash + +# usage: bash examples/on_policy_distillation/run-qwen3-8B-opd.sh + +set -ex + +export PYTHONUNBUFFERED=1 + +NVLINK_COUNT=$(nvidia-smi topo -m 2>/dev/null | grep -o 'NV[0-9][0-9]*' | wc -l) +if [ "$NVLINK_COUNT" -gt 0 ]; then + HAS_NVLINK=1 +else + HAS_NVLINK=0 +fi +echo "HAS_NVLINK: $HAS_NVLINK (detected $NVLINK_COUNT NVLink references)" + +# Start the teacher model server +TEACHER_IP="127.0.0.1" +TEACHER_PORT=13141 +LOG_FILE="/tmp/vllm_teacher_$(head /dev/urandom | tr -dc A-Za-z0-9 | head -c 6).log" + +## Launch the teacher model server in the background +CUDA_VISIBLE_DEVICES=4,5,6,7 python3 -m vllm.entrypoints.openai.api_server \ + --model /root/Qwen3-32B \ + --host 0.0.0.0 \ + --port ${TEACHER_PORT} \ + --tensor-parallel-size 4 \ + --gpu-memory-utilization 0.85 \ + --trust-remote-code \ + --dtype bfloat16 \ + --max-model-len 16384 \ + --disable-custom-all-reduce \ + > "${LOG_FILE}" 2>&1 & +TEACHER_PID=$! + +echo "Starting teacher model server (pid=${TEACHER_PID})..." + +## Wait for the teacher model server to be ready +for i in $(seq 1 120); do + if ! kill -0 "${TEACHER_PID}" 2>/dev/null; then + echo "ERROR: Teacher server process died. Check ${LOG_FILE}" + tail -n 20 "${LOG_FILE}" + exit 1 + fi + HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" "http://${TEACHER_IP}:${TEACHER_PORT}/health" 2>/dev/null || true) + if [ "${HTTP_CODE}" = "200" ]; then + echo "Teacher model server is up and running at ${TEACHER_IP}:${TEACHER_PORT}." + break + fi + if [ "$i" -eq 120 ]; then + echo "ERROR: Teacher server failed to start within 10 minutes" + tail -n 20 "${LOG_FILE}" + kill "${TEACHER_PID}" 2>/dev/null || true + exit 1 + fi + echo "Waiting for the teacher model server to start..." + sleep 5 +done +sleep 5 + +source "/root/vime/scripts/models/qwen3-8B.sh" + +CKPT_ARGS=( + --hf-checkpoint /root/Qwen3-8B + --ref-load /root/Qwen3-8B_torch_dist + --load /root/Qwen3-8B_torch_dist + --save /root/Qwen3-8B_vime/ + --save-interval 20 + --megatron-to-hf-mode bridge +) + +ROLLOUT_ARGS=( + --prompt-data /root/dapo-math-17k/dapo-math-17k.jsonl + --input-key prompt + --apply-chat-template + --rollout-shuffle + --rm-type math + --num-rollout 300 + --rollout-batch-size 16 + --n-samples-per-prompt 4 + --rollout-max-response-len 4096 + --rollout-max-context-len 8192 + --rollout-temperature 1 + + --global-batch-size 64 + --balance-data +) + +RM_ARGS=( + --custom-rm-path vime.rollout.on_policy_distillation.reward_func + --custom-reward-post-process-path vime.rollout.on_policy_distillation.post_process_rewards + --rm-url http://${TEACHER_IP}:${TEACHER_PORT}/inference/v1/generate +) + +EVAL_ARGS=( + # --eval-interval 50 + # --eval-prompt-data gsm8k /root/gsm8k/test.parquet + # --eval-input-key messages + # --n-samples-per-eval-prompt 1 + # --eval-max-response-len 4096 + # --eval-top-k 1 +) + +PERF_ARGS=( + --tensor-model-parallel-size 2 + --sequence-parallel + --pipeline-model-parallel-size 1 + --context-parallel-size 1 + --expert-model-parallel-size 1 + --expert-tensor-parallel-size 1 + + --recompute-granularity full + --recompute-method uniform + --recompute-num-layers 1 + + --use-dynamic-batch-size + --max-tokens-per-gpu 2048 +) + +GRPO_ARGS=( + --advantage-estimator grpo + --use-opd + --opd-type vllm + --opd-kl-coef 1.0 + --use-kl-loss + --kl-loss-coef 0.00 + --kl-loss-type low_var_kl + --entropy-coef 0.00 + --eps-clip 0.2 + --eps-clip-high 0.28 +) + +OPTIMIZER_ARGS=( + --optimizer adam + --lr 1e-6 + --lr-decay-style constant + --weight-decay 0.1 + --adam-beta1 0.9 + --adam-beta2 0.98 +) + +WANDB_ARGS=( + #--use-wandb + # --wandb-project vime-opd + # --wandb-group qwen3-8B-opd + # --wandb-key ${WANDB_KEY} +) + +VLLM_ARGS=( + --rollout-num-gpus-per-engine 1 + --vllm-gpu-memory-utilization 0.25 +) + +MISC_ARGS=( + --attention-dropout 0.0 + --hidden-dropout 0.0 + --accumulate-allreduce-grads-in-fp32 + --attention-softmax-in-fp32 + --attention-backend flash + --make-vocab-size-divisible-by 128 +) + +# launch the master node of ray in container +export CUDA_VISIBLE_DEVICES=0,1,2,3 +export MASTER_ADDR=${MASTER_ADDR:-"127.0.0.1"} +ray start --head --node-ip-address ${MASTER_ADDR} --num-gpus 4 --disable-usage-stats --dashboard-host=0.0.0.0 --dashboard-port=8265 + +RUNTIME_ENV_JSON="{ + \"env_vars\": { + \"PYTHONPATH\": \"/root/vime:/root/Megatron-LM/\", + \"CUDA_DEVICE_MAX_CONNECTIONS\": \"1\", + \"NCCL_NVLS_ENABLE\": \"${HAS_NVLINK}\" + } +}" + +ray job submit --address="http://127.0.0.1:8265" \ + --runtime-env-json="${RUNTIME_ENV_JSON}" \ + --working-dir /root/vime \ + -- python3 train.py \ + --train-backend megatron \ + --actor-num-nodes 1 \ + --actor-num-gpus-per-node 4 \ + --colocate \ + ${MODEL_ARGS[@]} \ + ${CKPT_ARGS[@]} \ + ${ROLLOUT_ARGS[@]} \ + ${OPTIMIZER_ARGS[@]} \ + ${GRPO_ARGS[@]} \ + ${WANDB_ARGS[@]} \ + ${PERF_ARGS[@]} \ + ${EVAL_ARGS[@]} \ + ${VLLM_ARGS[@]} \ + ${MISC_ARGS[@]} \ + ${RM_ARGS[@]} + +#### clear after training +kill ${TEACHER_PID} 2>/dev/null || true +sleep 3 +ray stop --force +pkill -9 ray +pkill -9 -f "train.py" || true +sleep 3