Files
myLightningOPD/scripts/run_poe_distill_qwen3_4b_lora.sh
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

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 \
"$@"