J-space / scripts /fit_multi_gpu.sh
ayh015's picture
Upload folder using huggingface_hub
f6158c7 verified
Raw
History Blame Contribute Delete
2.57 kB
#!/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"