#!/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"