127 lines
4.8 KiB
Bash
127 lines
4.8 KiB
Bash
#!/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 \
|
|
"$@"
|