File size: 2,571 Bytes
85b17bd
 
 
 
 
 
 
f6158c7
 
 
85b17bd
f6158c7
 
85b17bd
f6158c7
 
85b17bd
 
 
 
 
f6158c7
 
 
 
 
85b17bd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f6158c7
 
85b17bd
 
 
 
f6158c7
85b17bd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env bash
set -euo pipefail

# Examples:
#   GPUS=0,1 NUM_PROMPTS=20 bash scripts/fit_multi_gpu.sh
#   GPUS=0,1,2,3 NUM_PROMPTS=1000 OUTPUT_DIR=outputs/main-1000 bash scripts/fit_multi_gpu.sh

GPUS="${GPUS:-0,1}"
NUM_PROMPTS="${NUM_PROMPTS:-20}"
MODEL_PATH="${MODEL_PATH:-LLMs/qwen3-4b-base-sft-qwen3-8b}"
DATA_PATH="${DATA_PATH:-data/dapo-math-17k/dapo-math-17k.jsonl}"
CORPUS_FORMAT="${CORPUS_FORMAT:-auto}"
RESPONSE_WINDOW_LEN="${RESPONSE_WINDOW_LEN:-1024}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/multi-gpu-${NUM_PROMPTS}}"
DIM_BATCH="${DIM_BATCH:-8}"
MAX_SEQ_LEN="${MAX_SEQ_LEN:-}"
SKIP_FIRST="${SKIP_FIRST:-16}"
CHECKPOINT_EVERY="${CHECKPOINT_EVERY:-10}"
SEED="${SEED:-17}"
PYTHON_BIN="${PYTHON_BIN:-python}"

MAX_SEQ_ARGS=()
if [[ -n "${MAX_SEQ_LEN}" ]]; then
  MAX_SEQ_ARGS=(--max-seq-len "${MAX_SEQ_LEN}")
fi

IFS=',' read -r -a GPU_LIST <<< "${GPUS}"
NUM_GPUS="${#GPU_LIST[@]}"
if (( NUM_GPUS == 0 )); then
  echo "GPUS must contain at least one CUDA device ID" >&2
  exit 2
fi
if (( NUM_PROMPTS < NUM_GPUS )); then
  echo "NUM_PROMPTS (${NUM_PROMPTS}) must be at least the GPU count (${NUM_GPUS})" >&2
  exit 2
fi

mkdir -p "${OUTPUT_DIR}"
PIDS=()
CHECKPOINTS=()
OFFSET=0
BASE_COUNT=$((NUM_PROMPTS / NUM_GPUS))
REMAINDER=$((NUM_PROMPTS % NUM_GPUS))

for INDEX in "${!GPU_LIST[@]}"; do
  GPU="${GPU_LIST[$INDEX]}"
  COUNT="${BASE_COUNT}"
  if (( INDEX < REMAINDER )); then
    COUNT=$((COUNT + 1))
  fi
  SHARD_DIR="${OUTPUT_DIR}/shard-${INDEX}"
  CHECKPOINTS+=("${SHARD_DIR}/fit-checkpoint-fp32.pt")
  echo "Launching shard ${INDEX}: physical GPU ${GPU}, offset ${OFFSET}, count ${COUNT}"
  CUDA_VISIBLE_DEVICES="${GPU}" "${PYTHON_BIN}" -m math_jlens.cli \
    --model "${MODEL_PATH}" \
    --data "${DATA_PATH}" \
    --corpus-format "${CORPUS_FORMAT}" \
    --response-window-len "${RESPONSE_WINDOW_LEN}" \
    --output-dir "${SHARD_DIR}" \
    --num-prompts "${COUNT}" \
    --offset "${OFFSET}" \
    --seed "${SEED}" \
    "${MAX_SEQ_ARGS[@]}" \
    --skip-first "${SKIP_FIRST}" \
    --dim-batch "${DIM_BATCH}" \
    --checkpoint-every "${CHECKPOINT_EVERY}" \
    --device cuda:0 &
  PIDS+=("$!")
  OFFSET=$((OFFSET + COUNT))
done

FAILED=0
for INDEX in "${!PIDS[@]}"; do
  if ! wait "${PIDS[$INDEX]}"; then
    echo "Shard ${INDEX} failed" >&2
    FAILED=1
  fi
done
if (( FAILED != 0 )); then
  echo "At least one shard failed; partial checkpoints were kept for resume" >&2
  exit 1
fi

"${PYTHON_BIN}" -m math_jlens.merge \
  --checkpoints "${CHECKPOINTS[@]}" \
  --output "${OUTPUT_DIR}/lens-bf16.pt"

echo "Finished: ${OUTPUT_DIR}/lens-bf16.pt"