初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
42
scripts/collect_rollouts.sh
Normal file
42
scripts/collect_rollouts.sh
Normal file
@@ -0,0 +1,42 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Step 3: Collect student rollouts on OPD prompts.
|
||||
#
|
||||
# Uses data_curation/ to run the SFT model on OPD prompts (e.g. DAPO-Math-17k)
|
||||
# and collect response rollouts for Lightning OPD data preparation.
|
||||
#
|
||||
# Required environment variables:
|
||||
# SFT_CHECKPOINT - Path to the SFT model checkpoint
|
||||
# OPD_PROMPTS - Path to the OPD prompt dataset (.jsonl or .parquet)
|
||||
# OUTPUT_DIR - Directory for collected rollout data
|
||||
#
|
||||
# Optional:
|
||||
# NUM_GPUS - Number of GPUs to use (default: 8)
|
||||
# TP_SIZE - Tensor parallel size per worker (default: 1)
|
||||
#
|
||||
# Extra args are passed through to data_curation/pipeline.py, e.g.:
|
||||
# bash scripts/collect_rollouts.sh --num-samples 10
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
: "${SFT_CHECKPOINT:?Set SFT_CHECKPOINT to the SFT model path}"
|
||||
: "${OPD_PROMPTS:?Set OPD_PROMPTS to the OPD prompt dataset path}"
|
||||
: "${OUTPUT_DIR:?Set OUTPUT_DIR for collected rollout data}"
|
||||
|
||||
# Resolve to absolute paths (workers may run from different cwd)
|
||||
SFT_CHECKPOINT="$(cd "$(dirname "${SFT_CHECKPOINT}")" && pwd)/$(basename "${SFT_CHECKPOINT}")"
|
||||
OPD_PROMPTS="$(cd "$(dirname "${OPD_PROMPTS}")" && pwd)/$(basename "${OPD_PROMPTS}")"
|
||||
OUTPUT_DIR="$(mkdir -p "${OUTPUT_DIR}" && cd "${OUTPUT_DIR}" && pwd)"
|
||||
|
||||
NUM_GPUS="${NUM_GPUS:-8}"
|
||||
TP_SIZE="${TP_SIZE:-1}"
|
||||
|
||||
bash data_curation/run_curation.sh \
|
||||
--model "${SFT_CHECKPOINT}" \
|
||||
--input "${OPD_PROMPTS}" \
|
||||
--output-dir "${OUTPUT_DIR}" \
|
||||
--num-gpus "${NUM_GPUS}" \
|
||||
--tensor-parallel-size "${TP_SIZE}" \
|
||||
"$@"
|
||||
27
scripts/convert_megatron_to_hf.sh
Normal file
27
scripts/convert_megatron_to_hf.sh
Normal file
@@ -0,0 +1,27 @@
|
||||
#!/bin/bash
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# Convert a Megatron torch_dist checkpoint to HuggingFace format.
|
||||
#
|
||||
# Required env vars:
|
||||
# MEGATRON_CKPT_DIR - path to the Megatron checkpoint directory (e.g., /root/models/<name>_ckpt__<config>/iter_0000150)
|
||||
# HF_OUTPUT_DIR - path to save the converted HuggingFace model
|
||||
# ORIGIN_HF_DIR - path to the original HuggingFace model (for config.json, tokenizer, etc.)
|
||||
#
|
||||
# Example:
|
||||
# MEGATRON_CKPT_DIR=/root/models/Qwen3-4B-Base-sft_ckpt__qwen3-4b-lightning-opd/iter_0000150 \
|
||||
# HF_OUTPUT_DIR=checkpoints/qwen3-4b-lightning-opd-hf \
|
||||
# ORIGIN_HF_DIR=checkpoints/qwen3-4b-base-sft-qwen3-8b/<your-sft-checkpoint> \
|
||||
# bash scripts/convert_megatron_to_hf.sh
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
: "${MEGATRON_CKPT_DIR:?Please set MEGATRON_CKPT_DIR}"
|
||||
: "${HF_OUTPUT_DIR:?Please set HF_OUTPUT_DIR}"
|
||||
: "${ORIGIN_HF_DIR:?Please set ORIGIN_HF_DIR}"
|
||||
|
||||
python tools/convert_torch_dist_to_hf.py \
|
||||
--input-dir "${MEGATRON_CKPT_DIR}" \
|
||||
--output-dir "${HF_OUTPUT_DIR}" \
|
||||
--origin-hf-dir "${ORIGIN_HF_DIR}"
|
||||
53
scripts/eval_aime2024.sh
Normal file
53
scripts/eval_aime2024.sh
Normal file
@@ -0,0 +1,53 @@
|
||||
export MKL_THREADING_LAYER=GNU
|
||||
export MKL_SERVICE_FORCE_INTEL=0
|
||||
export OMP_NUM_THREADS=1
|
||||
|
||||
# CUDA_VISIBLE_DEVICES=0 python tools/eval_aime2024_vllm.py \
|
||||
# --model /mnt/disk1/yihao/Lightning-OPD/model_weights/qwen3-8b \
|
||||
# --num-gpus 1 \
|
||||
# --n-samples 1 \
|
||||
# --temperature 0.0 \
|
||||
# --top-p 1.0 \
|
||||
# --max-tokens 32768 \
|
||||
# --prompt-template paper \
|
||||
# --hf-cache /mnt/disk1/yihao/hf_cache \
|
||||
# --output outputs/aime2024_qwen3_8b_paper_n1_32k.jsonl
|
||||
|
||||
# CUDA_VISIBLE_DEVICES=0 python tools/eval_aime2024_vllm.py \
|
||||
# --model /mnt/disk1/yihao/Lightning-OPD/checkpoints/qwen3-4b-poe-distill-lora-opd-ppo-clip-locking-b-5-self-distill-100 \
|
||||
# --num-gpus 1 \
|
||||
# --n-samples 1 \
|
||||
# --temperature 0.0 \
|
||||
# --top-p 1.0 \
|
||||
# --max-tokens 32768 \
|
||||
# --prompt-template paper \
|
||||
# --hf-cache /mnt/disk1/yihao/hf_cache \
|
||||
# --output outputs/aime2024_qwen3_8b_paper_n1_32k.jsonl
|
||||
|
||||
|
||||
CUDA_VISIBLE_DEVICES=4,5,6,7 python tools/eval_aime2024_vllm.py \
|
||||
--model /mnt/disk1/yihao/Lightning-OPD/checkpoints/qwen3-4b-poe-distill-lora-opd-ppo-clip-locking-b-5-self-distill-100 \
|
||||
--num-gpus 4 \
|
||||
--n-samples 32 \
|
||||
--temperature 0.6 \
|
||||
--top-p 0.95 \
|
||||
--max-tokens 32768 \
|
||||
--prompt-template paper \
|
||||
--hf-cache /mnt/disk1/yihao/hf_cache \
|
||||
--output outputs/aime2024_qwen3_4b_poe_distill_lora_paper_n1_32k.jsonl \
|
||||
--enable-thinking
|
||||
|
||||
|
||||
|
||||
# CUDA_VISIBLE_DEVICES=4,5,6,7 python tools/eval_aime2024_vllm.py \
|
||||
# --model /mnt/disk1/yihao/Lightning-OPD/checkpoints/qwen3-4b-poe-distill-lora-opd-ppo-clip-60 \
|
||||
# --num-gpus 4 \
|
||||
# --n-samples 32 \
|
||||
# --temperature 0.6 \
|
||||
# --top-p 0.95 \
|
||||
# --max-tokens 32768 \
|
||||
# --enable-thinking \
|
||||
# --prompt-template paper \
|
||||
# --hf-cache /mnt/disk1/yihao/hf_cache \
|
||||
# --output outputs/aime2024_qwen3_4b_lightning_opd_paper_n32_32k.jsonl \
|
||||
|
||||
23
scripts/eval_aime2025.sh
Normal file
23
scripts/eval_aime2025.sh
Normal file
@@ -0,0 +1,23 @@
|
||||
CUDA_VISIBLE_DEVICES=4,5,6,7 python tools/eval_aime2025_vllm.py \
|
||||
--model /mnt/disk1/yihao/Lightning-OPD/checkpoints/qwen3-4b-lightning-opd-hf \
|
||||
--num-gpus 4 \
|
||||
--prompt-template paper \
|
||||
--n-samples 32 \
|
||||
--temperature 0.6 \
|
||||
--top-p 0.95 \
|
||||
--max-tokens 32768 \
|
||||
--enable-thinking \
|
||||
--output outputs/aime2025_qwen3_4b_lightning_opd.jsonl
|
||||
|
||||
|
||||
# CUDA_VISIBLE_DEVICES=0,2 python tools/eval_aime2025_vllm.py \
|
||||
# --model /mnt/disk1/yihao/Lightning-OPD/checkpoints/qwen3-4b-lightning-opd-hf \
|
||||
# --num-gpus 2 \
|
||||
# --n-samples 1 \
|
||||
# --temperature 0.0 \
|
||||
# --top-p 1.0 \
|
||||
# --max-tokens 32768 \
|
||||
# --prompt-template paper \
|
||||
# --hf-cache /mnt/disk1/yihao/hf_cache \
|
||||
# --output outputs/aime2024_qwen3_4b_poe_distill_lora_paper_n1_32k.jsonl \
|
||||
# --enable-thinking
|
||||
10
scripts/eval_hmmt25.sh
Normal file
10
scripts/eval_hmmt25.sh
Normal file
@@ -0,0 +1,10 @@
|
||||
CUDA_VISIBLE_DEVICES=4,5,6,7 python tools/eval_hmmt2025_vllm.py \
|
||||
--model /mnt/disk1/yihao/Lightning-OPD/checkpoints/qwen3-4b-lightning-opd-hf \
|
||||
--num-gpus 4 \
|
||||
--prompt-template paper \
|
||||
--n-samples 32 \
|
||||
--temperature 0.6 \
|
||||
--top-p 0.95 \
|
||||
--max-tokens 32768 \
|
||||
--enable-thinking \
|
||||
--output outputs/hmmt_feb_2025_qwen3_4b_lightning_opd.jsonl
|
||||
59
scripts/generate_sft_data.sh
Normal file
59
scripts/generate_sft_data.sh
Normal file
@@ -0,0 +1,59 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Step 1: Generate SFT data using the teacher model.
|
||||
#
|
||||
# Uses data_curation/ to run the teacher model on OpenThoughts-3 prompts
|
||||
# and generate response trajectories for SFT training.
|
||||
#
|
||||
# Required environment variables:
|
||||
# TEACHER_MODEL - HuggingFace model name or path (e.g. Qwen/Qwen3-8B)
|
||||
# SFT_PROMPTS - Path to the prompt dataset (.jsonl or .parquet)
|
||||
# OUTPUT_DIR - Directory for generated SFT data
|
||||
#
|
||||
# Optional:
|
||||
# NUM_GPUS - Number of GPUs to use (default: 8)
|
||||
# TP_SIZE - Tensor parallel size per worker (default: 1)
|
||||
#
|
||||
# Extra args are passed through to data_curation/pipeline.py, e.g.:
|
||||
# bash scripts/generate_sft_data.sh --num-samples 10
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# ---- CUDA / FlashInfer build environment ----
|
||||
if [ -n "${CONDA_PREFIX:-}" ]; then
|
||||
export CUDA_HOME="${CUDA_HOME:-$CONDA_PREFIX}"
|
||||
export CUDA_PATH="${CUDA_PATH:-$CONDA_PREFIX}"
|
||||
export CUDACXX="${CUDACXX:-$CONDA_PREFIX/bin/nvcc}"
|
||||
|
||||
export PATH="$CONDA_PREFIX/bin:$PATH"
|
||||
export LD_LIBRARY_PATH="$CONDA_PREFIX/lib:$CONDA_PREFIX/lib64:/usr/lib/x86_64-linux-gnu:${LD_LIBRARY_PATH:-}"
|
||||
export LIBRARY_PATH="/usr/lib/x86_64-linux-gnu:${LIBRARY_PATH:-}"
|
||||
fi
|
||||
|
||||
echo "Using nvcc: $(which nvcc)"
|
||||
nvcc --version || true
|
||||
echo "CUDA_HOME=${CUDA_HOME:-}"
|
||||
echo "CUDACXX=${CUDACXX:-}"
|
||||
echo "LIBRARY_PATH=${LIBRARY_PATH:-}"
|
||||
# ---------------------------------------------
|
||||
|
||||
: "${TEACHER_MODEL:?Set TEACHER_MODEL (e.g. Qwen/Qwen3-8B)}"
|
||||
: "${SFT_PROMPTS:?Set SFT_PROMPTS to the prompt dataset path}"
|
||||
: "${OUTPUT_DIR:?Set OUTPUT_DIR for generated SFT data}"
|
||||
|
||||
# Resolve to absolute paths (workers may run from different cwd)
|
||||
SFT_PROMPTS="$(cd "$(dirname "${SFT_PROMPTS}")" && pwd)/$(basename "${SFT_PROMPTS}")"
|
||||
OUTPUT_DIR="$(mkdir -p "${OUTPUT_DIR}" && cd "${OUTPUT_DIR}" && pwd)"
|
||||
|
||||
NUM_GPUS="${NUM_GPUS:-8}"
|
||||
TP_SIZE="${TP_SIZE:-1}"
|
||||
|
||||
bash data_curation/run_curation.sh \
|
||||
--model "${TEACHER_MODEL}" \
|
||||
--input "${SFT_PROMPTS}" \
|
||||
--output-dir "${OUTPUT_DIR}" \
|
||||
--num-gpus "${NUM_GPUS}" \
|
||||
--tensor-parallel-size "${TP_SIZE}" \
|
||||
"$@"
|
||||
5
scripts/merge_poe_lora.sh
Normal file
5
scripts/merge_poe_lora.sh
Normal file
@@ -0,0 +1,5 @@
|
||||
python tools/merge_poe_lora.py \
|
||||
--base-model checkpoints/qwen3-4b-base-sft-qwen3-8b \
|
||||
--adapter checkpoints/qwen3-4b-poe-distill-lora-opd-ppo-clip-locking-b-5-self-distill/checkpoint-100 \
|
||||
--output-dir checkpoints/qwen3-4b-poe-distill-lora-opd-ppo-clip-locking-b-5-self-distill-100 \
|
||||
--dtype bfloat16
|
||||
29
scripts/precompute_teacher_logprobs_4b.sh
Normal file
29
scripts/precompute_teacher_logprobs_4b.sh
Normal file
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Precompute teacher logprobs for Lightning OPD (4B scale, teacher=Qwen3-8B).
|
||||
#
|
||||
# Required environment variables:
|
||||
# SFT_CHECKPOINT - Path to the SFT checkpoint (used as tokenizer)
|
||||
# ROLLOUT_PARQUET - Path to the student rollout parquet file
|
||||
# OUTPUT_DIR - Directory for the output parquet with teacher logprobs
|
||||
#
|
||||
# This script starts a Qwen3-8B teacher server, then runs Phase 1+2 of
|
||||
# prepare_lightning_opd.py to tokenize and precompute teacher logprobs.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
: "${SFT_CHECKPOINT:?Set SFT_CHECKPOINT to the SFT model path}"
|
||||
: "${ROLLOUT_PARQUET:?Set ROLLOUT_PARQUET to the student rollout parquet}"
|
||||
: "${OUTPUT_DIR:?Set OUTPUT_DIR for the output parquet}"
|
||||
|
||||
# Start teacher server
|
||||
bash scripts/serve_teacher_8b.sh
|
||||
|
||||
python3 data_curation/prepare_lightning_opd.py \
|
||||
--tokenizer-path "${SFT_CHECKPOINT}" \
|
||||
--input-parquet "${ROLLOUT_PARQUET}" \
|
||||
--output-dir "${OUTPUT_DIR}" \
|
||||
--compute-teacher-logprobs \
|
||||
--teacher-url http://127.0.0.1:13141/generate
|
||||
29
scripts/precompute_teacher_logprobs_8b.sh
Normal file
29
scripts/precompute_teacher_logprobs_8b.sh
Normal file
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Precompute teacher logprobs for Lightning OPD (8B scale, teacher=Qwen3-32B).
|
||||
#
|
||||
# Required environment variables:
|
||||
# SFT_CHECKPOINT - Path to the SFT checkpoint (used as tokenizer)
|
||||
# ROLLOUT_PARQUET - Path to the student rollout parquet file
|
||||
# OUTPUT_DIR - Directory for the output parquet with teacher logprobs
|
||||
#
|
||||
# This script starts a Qwen3-32B teacher server, then runs Phase 1+2 of
|
||||
# prepare_lightning_opd.py to tokenize and precompute teacher logprobs.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
: "${SFT_CHECKPOINT:?Set SFT_CHECKPOINT to the SFT model path}"
|
||||
: "${ROLLOUT_PARQUET:?Set ROLLOUT_PARQUET to the student rollout parquet}"
|
||||
: "${OUTPUT_DIR:?Set OUTPUT_DIR for the output parquet}"
|
||||
|
||||
# Start teacher server
|
||||
bash scripts/serve_teacher_32b.sh
|
||||
|
||||
python3 data_curation/prepare_lightning_opd.py \
|
||||
--tokenizer-path "${SFT_CHECKPOINT}" \
|
||||
--input-parquet "${ROLLOUT_PARQUET}" \
|
||||
--output-dir "${OUTPUT_DIR}" \
|
||||
--compute-teacher-logprobs \
|
||||
--teacher-url http://127.0.0.1:13141/generate
|
||||
150
scripts/prepare_sft_prompts.py
Normal file
150
scripts/prepare_sft_prompts.py
Normal file
@@ -0,0 +1,150 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
"""
|
||||
Convert HuggingFace OpenThoughts3-1.2M dataset to a prompt-only JSONL file
|
||||
for SFT data generation (Step 1).
|
||||
|
||||
Extracts the prompt (user messages) from each sample and writes to JSONL.
|
||||
Optionally samples a subset (default 300K) to reduce compute cost.
|
||||
|
||||
Usage:
|
||||
python scripts/prepare_sft_prompts.py \
|
||||
--output data/prompts/openthoughts3_300k.jsonl \
|
||||
--num-samples 300000
|
||||
|
||||
# Use a local parquet file instead of downloading from HF
|
||||
python scripts/prepare_sft_prompts.py \
|
||||
--input data/prompts/local.parquet \
|
||||
--output data/prompts/openthoughts3_300k.jsonl
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import random
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Extract prompts from OpenThoughts3-1.2M for SFT data generation."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--input", type=str, default=None,
|
||||
help="Path to a local parquet/jsonl file. If not set, downloads from HuggingFace.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--hf-dataset", type=str, default="open-thoughts/OpenThoughts3-1.2M",
|
||||
help="HuggingFace dataset name (default: open-thoughts/OpenThoughts3-1.2M).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output", type=str, required=True,
|
||||
help="Output JSONL file path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-samples", type=int, default=300000,
|
||||
help="Number of samples to keep (default: 300000). Set to 0 for all.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed", type=int, default=42,
|
||||
help="Random seed for sampling (default: 42).",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def extract_prompt(sample):
|
||||
"""Extract the prompt (non-assistant messages) from a sample.
|
||||
|
||||
Supports two common formats:
|
||||
1. {"conversations": [{"from": "human", "value": ...}, ...]} (sharegpt)
|
||||
2. {"prompt": [{"role": "user", "content": ...}, ...]} (chat messages)
|
||||
"""
|
||||
if "conversations" in sample:
|
||||
messages = []
|
||||
for turn in sample["conversations"]:
|
||||
role = turn.get("from", turn.get("role", ""))
|
||||
content = turn.get("value", turn.get("content", ""))
|
||||
if role in ("human", "user"):
|
||||
messages.append({"role": "user", "content": content})
|
||||
elif role == "system":
|
||||
messages.append({"role": "system", "content": content})
|
||||
if messages:
|
||||
return {"prompt": messages}
|
||||
|
||||
if "prompt" in sample:
|
||||
if isinstance(sample["prompt"], list):
|
||||
return {"prompt": sample["prompt"]}
|
||||
elif isinstance(sample["prompt"], str):
|
||||
return {"prompt": [{"role": "user", "content": sample["prompt"]}]}
|
||||
|
||||
if "messages" in sample:
|
||||
messages = [
|
||||
{"role": m["role"], "content": m["content"]}
|
||||
for m in sample["messages"]
|
||||
if m["role"] != "assistant"
|
||||
]
|
||||
if messages:
|
||||
return {"prompt": messages}
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def load_dataset_from_hf(dataset_name):
|
||||
"""Load dataset from HuggingFace."""
|
||||
from datasets import load_dataset
|
||||
print(f"Loading dataset from HuggingFace: {dataset_name}")
|
||||
ds = load_dataset(dataset_name, split="train")
|
||||
return ds
|
||||
|
||||
|
||||
def load_dataset_from_file(path):
|
||||
"""Load dataset from local file (parquet or jsonl)."""
|
||||
import pandas as pd
|
||||
print(f"Loading dataset from local file: {path}")
|
||||
if path.endswith(".parquet"):
|
||||
df = pd.read_parquet(path)
|
||||
return df.to_dict("records")
|
||||
elif path.endswith(".jsonl"):
|
||||
with open(path) as f:
|
||||
return [json.loads(line) for line in f]
|
||||
else:
|
||||
raise ValueError(f"Unsupported format: {path}")
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
random.seed(args.seed)
|
||||
|
||||
# Load dataset
|
||||
if args.input:
|
||||
samples = load_dataset_from_file(args.input)
|
||||
else:
|
||||
samples = load_dataset_from_hf(args.hf_dataset)
|
||||
|
||||
print(f"Total samples: {len(samples)}")
|
||||
|
||||
# Sample subset
|
||||
if args.num_samples > 0 and args.num_samples < len(samples):
|
||||
indices = random.sample(range(len(samples)), args.num_samples)
|
||||
indices.sort()
|
||||
samples = [samples[i] for i in indices]
|
||||
print(f"Sampled {args.num_samples} samples")
|
||||
|
||||
# Extract prompts
|
||||
from tqdm import tqdm
|
||||
written = 0
|
||||
skipped = 0
|
||||
with open(args.output, "w") as f:
|
||||
for sample in tqdm(samples, desc="Extracting prompts"):
|
||||
prompt_item = extract_prompt(sample)
|
||||
if prompt_item and len(prompt_item["prompt"]) > 0:
|
||||
f.write(json.dumps(prompt_item) + "\n")
|
||||
written += 1
|
||||
else:
|
||||
skipped += 1
|
||||
|
||||
print(f"Written: {written}, Skipped: {skipped}")
|
||||
print(f"Output: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
126
scripts/run_poe_distill_qwen3_4b_lora.sh
Normal file
126
scripts/run_poe_distill_qwen3_4b_lora.sh
Normal file
@@ -0,0 +1,126 @@
|
||||
#!/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 \
|
||||
"$@"
|
||||
19
scripts/serve_teacher_32b.sh
Normal file
19
scripts/serve_teacher_32b.sh
Normal file
@@ -0,0 +1,19 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
LOG_FILE="/tmp/sglang_$(head /dev/urandom | tr -dc A-Za-z0-9 | head -c 6).log"
|
||||
python3 -m sglang.launch_server \
|
||||
--model-path Qwen/Qwen3-32B \
|
||||
--host 127.0.0.1 \
|
||||
--port 13141 \
|
||||
--tp 8 \
|
||||
--chunked-prefill-size 4096 \
|
||||
--mem-fraction-static 0.6 \
|
||||
--context-length 8192 \
|
||||
> "$LOG_FILE" 2>&1 &
|
||||
|
||||
until curl -sf http://127.0.0.1:13141/health_generate > /dev/null; do
|
||||
echo "Waiting for the teacher model server to start..."
|
||||
tail -n 10 "$LOG_FILE"
|
||||
sleep 5
|
||||
done
|
||||
19
scripts/serve_teacher_8b.sh
Normal file
19
scripts/serve_teacher_8b.sh
Normal file
@@ -0,0 +1,19 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
LOG_FILE="/tmp/sglang_$(head /dev/urandom | tr -dc A-Za-z0-9 | head -c 6).log"
|
||||
python3 -m sglang.launch_server \
|
||||
--model-path model_weights/qwen3-8b \
|
||||
--host 127.0.0.1 \
|
||||
--port 13141 \
|
||||
--tp 4 \
|
||||
--chunked-prefill-size 4096 \
|
||||
--mem-fraction-static 0.6 \
|
||||
--context-length 8192 \
|
||||
> "$LOG_FILE" 2>&1 &
|
||||
|
||||
until curl -sf http://127.0.0.1:13141/health_generate > /dev/null; do
|
||||
echo "Waiting for the teacher model server to start..."
|
||||
tail -n 10 "$LOG_FILE"
|
||||
sleep 5
|
||||
done
|
||||
Reference in New Issue
Block a user