| #!/usr/bin/env bash |
| set -euo pipefail |
|
|
| |
| |
| |
|
|
| 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" |
|
|