#!/usr/bin/env bash # SPDX-License-Identifier: Apache-2.0 # Starter launcher for PoE / sampled-token OPD LoRA training with: # 1) hold-then-transition beta schedule # 2) optional hold-then-transition learning-rate schedule # # Default schedule below: # steps 0-99: beta = 1.0, lr = 2e-6 # steps 100-110: beta 1.0 -> 0.5, lr 2e-6 -> 1e-7 # steps 111+: beta = 0.5, lr = 1e-7 # # Note: these are optimizer global steps, not micro-batch steps. set -euo pipefail : "${SFT_CHECKPOINT:?Set SFT_CHECKPOINT to the Qwen3-4B SFT checkpoint}" # Path to the patched Python trainer. If you copy the Python file into tools/, # leave this default; otherwise pass SCRIPT_PATH=/path/to/file.py. SCRIPT_PATH="${SCRIPT_PATH:-tools/train_poe_distill_lora.py}" TEACHER_MODEL="${TEACHER_MODEL:-model_weights/qwen3-8b}" TRAIN_DATA="${TRAIN_DATA:-data/rollouts/dapo-math-17k-qwen3-4b-sft-rollouts.parquet}" OUTPUT_DIR="${OUTPUT_DIR:-checkpoints/qwen3-4b-poe-distill-lora-opd-hold-beta-lr}" CUDA_VISIBLE_DEVICES="${CUDA_VISIBLE_DEVICES:-0}" NPROC_PER_NODE="${NPROC_PER_NODE:-1}" # Training length. MAX_STEPS="${MAX_STEPS:-120}" SAVE_STEPS="${SAVE_STEPS:-10}" LOGGING_STEPS="${LOGGING_STEPS:-1}" # Beta schedule: hold beta_start, transition to beta_end. BETA_START="${BETA_START:-1.0}" BETA_END="${BETA_END:-0.5}" BETA_HOLD_STEPS="${BETA_HOLD_STEPS:-100}" BETA_TRANSITION_STEPS="${BETA_TRANSITION_STEPS:-10}" HOLD_TRANSITION_SCHEDULE="${HOLD_TRANSITION_SCHEDULE:-linear}" # linear or cosine # LR schedule: hold lr_start, transition to lr_end. # When LR_END is non-empty, use constant HF scheduler and let the callback set LR. LEARNING_RATE="${LEARNING_RATE:-2e-6}" LR_START="${LR_START:-${LEARNING_RATE}}" LR_END="${LR_END:-1e-7}" LR_HOLD_STEPS="${LR_HOLD_STEPS:-${BETA_HOLD_STEPS}}" LR_TRANSITION_STEPS="${LR_TRANSITION_STEPS:-${BETA_TRANSITION_STEPS}}" # Loss / OPD behavior. LOSS_TYPE="${LOSS_TYPE:-sampled_token}" ADVANTAGE_NORMALIZATION="${ADVANTAGE_NORMALIZATION:-none}" ADVANTAGE_CLIP="${ADVANTAGE_CLIP:-10.0}" USE_PPO_CLIP="${USE_PPO_CLIP:-1}" PPO_CLIP_LOW="${PPO_CLIP_LOW:-0.2}" PPO_CLIP_HIGH="${PPO_CLIP_HIGH:-0.2}" SAMPLED_LOSS_REDUCTION="${SAMPLED_LOSS_REDUCTION:-per_sample}" POSITIVE_ADVANTAGES_ONLY="${POSITIVE_ADVANTAGES_ONLY:-0}" # Model / optimizer defaults. MAX_LENGTH="${MAX_LENGTH:-4096}" DISTILL_CHUNK_SIZE="${DISTILL_CHUNK_SIZE:-128}" PER_DEVICE_TRAIN_BATCH_SIZE="${PER_DEVICE_TRAIN_BATCH_SIZE:-2}" GRADIENT_ACCUMULATION_STEPS="${GRADIENT_ACCUMULATION_STEPS:-8}" WEIGHT_DECAY="${WEIGHT_DECAY:-0.1}" ADAM_BETA1="${ADAM_BETA1:-0.9}" ADAM_BETA2="${ADAM_BETA2:-0.98}" WARMUP_RATIO="${WARMUP_RATIO:-0.0}" LR_SCHEDULER_TYPE="${LR_SCHEDULER_TYPE:-constant}" FREEZE_LORA_B_AFTER_STEP="${FREEZE_LORA_B_AFTER_STEP:-99999}" PPO_CLIP_ARGS=() if [[ "${USE_PPO_CLIP}" == "1" || "${USE_PPO_CLIP}" == "true" || "${USE_PPO_CLIP}" == "True" ]]; then PPO_CLIP_ARGS+=(--use-ppo-clip) fi POS_ADV_ARGS=() if [[ "${POSITIVE_ADVANTAGES_ONLY}" == "1" || "${POSITIVE_ADVANTAGES_ONLY}" == "true" || "${POSITIVE_ADVANTAGES_ONLY}" == "True" ]]; then POS_ADV_ARGS+=(--positive-advantages-only) fi LR_ARGS=(--lr-start "${LR_START}") if [[ -n "${LR_END}" ]]; then LR_ARGS+=( --lr-end "${LR_END}" --lr-hold-steps "${LR_HOLD_STEPS}" --lr-transition-steps "${LR_TRANSITION_STEPS}" ) fi CUDA_VISIBLE_DEVICES="${CUDA_VISIBLE_DEVICES}" torchrun --standalone --nproc_per_node="${NPROC_PER_NODE}" "${SCRIPT_PATH}" \ --student-model "${SFT_CHECKPOINT}" \ --teacher-model "${TEACHER_MODEL}" \ --train-data "${TRAIN_DATA}" \ --output-dir "${OUTPUT_DIR}" \ --alpha "${ALPHA:-1.0}" \ --beta-start "${BETA_START}" \ --beta-end "${BETA_END}" \ --beta-hold-steps "${BETA_HOLD_STEPS}" \ --beta-transition-steps "${BETA_TRANSITION_STEPS}" \ --hold-transition-schedule "${HOLD_TRANSITION_SCHEDULE}" \ "${LR_ARGS[@]}" \ --loss-type "${LOSS_TYPE}" \ --advantage-normalization "${ADVANTAGE_NORMALIZATION}" \ --advantage-clip "${ADVANTAGE_CLIP}" \ "${PPO_CLIP_ARGS[@]}" \ --ppo-clip-low "${PPO_CLIP_LOW}" \ --ppo-clip-high "${PPO_CLIP_HIGH}" \ --sampled-loss-reduction "${SAMPLED_LOSS_REDUCTION}" \ "${POS_ADV_ARGS[@]}" \ --max-length "${MAX_LENGTH}" \ --distill-chunk-size "${DISTILL_CHUNK_SIZE}" \ --per-device-train-batch-size "${PER_DEVICE_TRAIN_BATCH_SIZE}" \ --gradient-accumulation-steps "${GRADIENT_ACCUMULATION_STEPS}" \ --learning-rate "${LEARNING_RATE}" \ --weight-decay "${WEIGHT_DECAY}" \ --adam-beta1 "${ADAM_BETA1}" \ --adam-beta2 "${ADAM_BETA2}" \ --warmup-ratio "${WARMUP_RATIO}" \ --lr-scheduler-type "${LR_SCHEDULER_TYPE}" \ --max-steps "${MAX_STEPS}" \ --save-steps "${SAVE_STEPS}" \ --logging-steps "${LOGGING_STEPS}" \ --freeze-lora-b-after-step "${FREEZE_LORA_B_AFTER_STEP}" \ --no-gradient-checkpointing \ "$@"