File size: 5,117 Bytes
58258b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
#!/usr/bin/env bash
# Qwen3-8B self-instill SFT on a single NVIDIA node. Local disk only: no GCS, no Iris.
#
#   ./launch_sft_8b.sh <hf-dataset-id> <7-char-revision> <run-name>
#
# Stage 1 (CPU, ~10-30 min) is cached; stage 2 is the training run.
set -euo pipefail

DATASET_ID="${1:?usage: launch_sft_8b.sh <dataset-id> <revision> <run-name>}"
REVISION="${2:?}"
RUN_NAME="${3:?}"

ROOT="${SFT_ROOT:-/data/sft}"
DATA_DIR="${ROOT}/data/$(echo "${DATASET_ID}" | tr '/' '_')-${REVISION}"
OUT_DIR="${ROOT}/runs/${RUN_NAME}"
NPROC="${NPROC:-$(nvidia-smi --list-gpus | wc -l)}"

export HF_HOME="${HF_HOME:-${ROOT}/hf}"
export TOKENIZERS_PARALLELISM=false
export OMP_NUM_THREADS="${OMP_NUM_THREADS:-8}"
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
export NCCL_DEBUG="${NCCL_DEBUG:-WARN}"
export PYTHONUNBUFFERED=1
# W&B is optional; the TPU runs used WANDB_MODE=offline on the workers.
export WANDB_MODE="${WANDB_MODE:-offline}"

mkdir -p "${OUT_DIR}"

# ----------------------------------------------------------------------------------
# Stage 1: HF dataset -> packed local .npy (idempotent, skipped if manifest exists)
# ----------------------------------------------------------------------------------
if [[ ! -f "${DATA_DIR}/manifest.json" ]]; then
  python prepare_sft_data.py \
    --dataset-id "${DATASET_ID}" \
    --revision "${REVISION}" \
    --tokenizer Qwen/Qwen3-8B \
    --max-seq-len 32768 \
    --text-replacements marin \
    --num-proc "$(nproc)" \
    --out "${DATA_DIR}" 2>&1 | tee "${OUT_DIR}/prepare.log"

  # one-time fidelity gate: our encoder must agree with levanter's segment-split encoder
  python prepare_sft_data.py \
    --dataset-id "${DATASET_ID}" --revision "${REVISION}" \
    --tokenizer Qwen/Qwen3-8B --verify 200 2>&1 | tee -a "${OUT_DIR}/prepare.log"
fi

# ----------------------------------------------------------------------------------
# Stage 1.5: print the plan (LR schedule, epochs, export cadence) before burning GPUs
# ----------------------------------------------------------------------------------
python train_sft_qwen3.py --data "${DATA_DIR}" --out "${OUT_DIR}" --recipe qwen3-8b --dry-run

# ----------------------------------------------------------------------------------
# Stage 2: training
#
# Batch geometry (must reproduce global batch 64 x 32768 tokens = 2,097,152 tok/step):
#   NPROC=8 : --per-device-batch 1, loss-groups 2 -> 4 micro-steps/group, 8 accum/step
#   NPROC=4 : --per-device-batch 1, loss-groups 2 -> 8 micro-steps/group, 16 accum/step
#   NPROC=16: --per-device-batch 1, loss-groups 2 -> 2 micro-steps/group, 4 accum/step
#
# Memory per GPU, Qwen3-8B (8.19e9 params), seq 32768, FSDP FULL_SHARD over 8 GPUs,
# fp32 master + bf16 compute + full activation checkpointing + chunked CE(2048):
#   fp32 master params      4 * 8.19e9 / 8      =  4.1 GB
#   Adam m + v (fp32)       8 * 8.19e9 / 8      =  8.2 GB
#   fp32 grad shard         4 * 8.19e9 / 8      =  4.1 GB
#   bf16 all-gather buffers root(1.24e9)+layers =  3.7 GB
#   checkpointed layer inputs 36*32768*4096*2   =  9.7 GB
#   one layer recompute working set             =  3.2 GB
#   chunked CE (2048 x 151936 fp32, fwd+grad)   =  2.5 GB
#   last hidden + grad                          =  0.5 GB
#   ----------------------------------------------------------
#   ~36 GB + NCCL/fragmentation  ->  ~40 GB      (fits 80 GB comfortably)
#
#   4 GPUs : optimiser/param/grad term triples to 32.8 GB -> ~50 GB, still fits 80 GB.
#   8x40 GB A100: ~30-34 GB, tight. Use --loss-chunk-size 1024 --wrap-embeddings,
#                 or 16 GPUs, or the 8K fallback below.
#
# Wall clock: ~3.9e20 FLOPs total (1.75e20 dense + 1.2e20 attention at 32K, x4/3 for
# recompute). 8xH100 at 30-35% MFU  ->  40-50 h. 8xB200/H200 ->  ~20-25 h.
# ----------------------------------------------------------------------------------
torchrun --standalone --nproc_per_node="${NPROC}" train_sft_qwen3.py \
  --data "${DATA_DIR}" \
  --out "${OUT_DIR}" \
  --recipe qwen3-8b \
  --train-batch-size 64 \
  --per-device-batch 1 \
  --loss-groups 2 \
  --num-train-steps 2000 \
  --steps-per-hf-export 100 \
  --steps-per-checkpoint 20 \
  --loss-chunk-size 2048 \
  --hf-export-dtype bfloat16 \
  --resume auto \
  2>&1 | tee -a "${OUT_DIR}/train.log"

# ----------------------------------------------------------------------------------
# 8K-context fallback (OOM at 32K, or no FlashAttention build).
# Keeps tokens/step, step count and LR schedule identical: 256 x 8192 = 2,097,152.
# Documents longer than 8192 are LEFT-SLICED (kept from the beginning) exactly as
# levanter's slice_strategy="left" does; check `docs_over_max_seq_len` in the
# prepare-stage stats before accepting this.
#
#   python prepare_sft_data.py ... --max-seq-len 8192 --out ${DATA_DIR}-8k
#   torchrun --standalone --nproc_per_node=8 train_sft_qwen3.py \
#     --data ${DATA_DIR}-8k --out ${OUT_DIR}-8k --recipe qwen3-8b \
#     --train-batch-size 256 --per-device-batch 4 --loss-groups 2 \
#     --num-train-steps 2000 --loss-chunk-size 2048
# ----------------------------------------------------------------------------------