|
| 1 | +#!/usr/bin/env bash |
| 2 | +# Copyright 2026 The LoongForge Authors. |
| 3 | +# SPDX-License-Identifier: Apache-2.0 |
| 4 | +# |
| 5 | +# FastWAM SFT with ZeRO Stage-1 (multi-GPU DDP). |
| 6 | +# |
| 7 | +# Delta versus run_fastwam_sft_ddp.sh: |
| 8 | +# --zero-optimizer wrap the optimizer in |
| 9 | +# ZeroRedundancyOptimizer, sharding |
| 10 | +# optimizer states across ranks. Only |
| 11 | +# effective with --distributed-strategy ddp. |
| 12 | +# --no-ddp-find-unused-parameters skip the unused-parameter scan; FastWAM's |
| 13 | +# forward graph has no conditional branches. |
| 14 | +# --ddp-static-graph graph is identical every iteration, lets |
| 15 | +# DDP reuse its bucket/reduction plan. |
| 16 | +# --ddp-gradient-as-bucket-view expose grads as views into the comm |
| 17 | +# buckets instead of separate allocations. |
| 18 | +# --no-ddp-broadcast-buffers no BN-style buffers to sync each forward. |
| 19 | +# --ddp-bucket-cap-mb larger buckets: fewer, bigger all-reduces. |
| 20 | +# |
| 21 | +# The memory saved by ZeRO-1 is what makes the larger --per-device-batch-size |
| 22 | +# below affordable relative to the plain DDP script. |
| 23 | +# |
| 24 | +# Two optional ZeRO knobs are left off by default: |
| 25 | +# --zero-parameters-as-bucket-view further cuts peak memory, but can clash |
| 26 | +# with torch.compile + the DDP reducer. |
| 27 | +# --zero-master-param-dtype fp32 rank-local fp32 master params, broadcast |
| 28 | +# after each step. Better numerics under |
| 29 | +# bf16 training at some bandwidth cost. |
| 30 | +# |
| 31 | +# Usage: |
| 32 | +# DATASET_PATH=/path/to/libero TOKENIZER_PATH=/path/to/tokenizer \ |
| 33 | +# bash run_fastwam_sft_zero1.sh |
| 34 | +# GPUS_PER_NODE=4 bash run_fastwam_sft_zero1.sh # override via env |
| 35 | +# ... bash run_fastwam_sft_zero1.sh --train-iters 50 # override a flag |
| 36 | + |
| 37 | +set -euo pipefail |
| 38 | + |
| 39 | +export LOONGFORGE_PATH="${LOONGFORGE_PATH:-/workspace/LoongForge}" |
| 40 | + |
| 41 | +DATASET_PATH=${DATASET_PATH:-/path/to/libero} |
| 42 | +TOKENIZER_PATH=${TOKENIZER_PATH:-/path/to/tokenizer} |
| 43 | +OUTPUT_DIR=${OUTPUT_DIR:-"outputs/fastwam_sft_zero1_$(date +%Y%m%d_%H%M%S)"} |
| 44 | + |
| 45 | +export CUBLAS_WORKSPACE_CONFIG=${CUBLAS_WORKSPACE_CONFIG:-:4096:8} |
| 46 | +export CUDA_DEVICE_MAX_CONNECTIONS=${CUDA_DEVICE_MAX_CONNECTIONS:-1} |
| 47 | + |
| 48 | +PRETRAINED_CHECKPOINT=${PRETRAINED_CHECKPOINT:-} |
| 49 | +ACTION_DIT_PRETRAINED_PATH=${ACTION_DIT_PRETRAINED_PATH:-} |
| 50 | +TEXT_EMBEDDING_CACHE_DIR=${TEXT_EMBEDDING_CACHE_DIR:-} |
| 51 | + |
| 52 | +GPUS_PER_NODE=${GPUS_PER_NODE:-8} |
| 53 | +MASTER_PORT=${MASTER_PORT:-29519} |
| 54 | + |
| 55 | +MODEL_NAME=${MODEL_NAME:-"fastwam"} |
| 56 | + |
| 57 | +TRAINING_ARGS=( |
| 58 | + --model-name "$MODEL_NAME" |
| 59 | + --trainer-type FinetuneTrainer |
| 60 | + --dtype bfloat16 |
| 61 | + --train-iters 20000 |
| 62 | + --save-interval 2000 |
| 63 | + --clip-grad 1.0 |
| 64 | + --gradient-accumulation-steps 1 |
| 65 | + --log-interval 1 |
| 66 | + --seed 3047 |
| 67 | + --output-dir "$OUTPUT_DIR" |
| 68 | + --tokenizer-path "$TOKENIZER_PATH" |
| 69 | +) |
| 70 | + |
| 71 | +# ── ZeRO Stage-1 + DDP tuning (the point of this script) ────── |
| 72 | +ZERO1_ARGS=( |
| 73 | + --distributed-strategy ddp |
| 74 | + --zero-optimizer |
| 75 | + --no-ddp-find-unused-parameters |
| 76 | + --ddp-static-graph |
| 77 | + --ddp-gradient-as-bucket-view |
| 78 | + --no-ddp-broadcast-buffers |
| 79 | + --ddp-bucket-cap-mb 200 |
| 80 | +) |
| 81 | + |
| 82 | +if [[ -n "$PRETRAINED_CHECKPOINT" ]]; then |
| 83 | + TRAINING_ARGS+=(--pretrained-checkpoint "$PRETRAINED_CHECKPOINT") |
| 84 | +fi |
| 85 | + |
| 86 | +LR_ARGS=( |
| 87 | + --lr-base 1.0e-8 |
| 88 | + --lr-decay-style cosine_warmup_with_min_lr |
| 89 | + --lr-warmup-iters 0 |
| 90 | + --min-lr 1.0e-9 |
| 91 | + --weight-decay 0.01 |
| 92 | + --adam-beta1 0.9 |
| 93 | + --adam-beta2 0.95 |
| 94 | +) |
| 95 | + |
| 96 | +DATA_ARGS=( |
| 97 | + --dataset-format lerobot_datasets |
| 98 | + --dataset-strategy fastwam |
| 99 | + --dataset-path "$DATASET_PATH" |
| 100 | + --robot-type libero_franka |
| 101 | + --per-device-batch-size 16 |
| 102 | + --num-workers 16 |
| 103 | + --lerobotdataset-version v2.1 |
| 104 | + --video-backend pyav |
| 105 | +) |
| 106 | + |
| 107 | +LOGGING_ARGS=( |
| 108 | + --wandb-project loongforge-vla |
| 109 | + --wandb-mode disabled |
| 110 | +) |
| 111 | + |
| 112 | +# ── Model/data dotlist overrides ────────────────────────────── |
| 113 | +MODEL_DATA_OVERRIDES=() |
| 114 | +if [[ -n "$ACTION_DIT_PRETRAINED_PATH" ]]; then |
| 115 | + MODEL_DATA_OVERRIDES+=("model.action_dit_pretrained_path=$ACTION_DIT_PRETRAINED_PATH") |
| 116 | +fi |
| 117 | +if [[ -n "$TEXT_EMBEDDING_CACHE_DIR" ]]; then |
| 118 | + MODEL_DATA_OVERRIDES+=("data.text_embedding_cache_dir=$TEXT_EMBEDDING_CACHE_DIR") |
| 119 | +fi |
| 120 | + |
| 121 | +mkdir -p "$OUTPUT_DIR" |
| 122 | + |
| 123 | +echo "════════════════════════════════════════════════════════════" |
| 124 | +echo " LoongForge FastWAM SFT (DDP + ZeRO-1)" |
| 125 | +echo " Model: $MODEL_NAME" |
| 126 | +echo " GPUs: $GPUS_PER_NODE" |
| 127 | +echo " Data: $DATASET_PATH" |
| 128 | +echo " Output: $OUTPUT_DIR" |
| 129 | +echo "════════════════════════════════════════════════════════════" |
| 130 | + |
| 131 | +PYTHONPATH="$LOONGFORGE_PATH:${PYTHONPATH:-}" \ |
| 132 | + torchrun --nproc_per_node "$GPUS_PER_NODE" --master_port "$MASTER_PORT" \ |
| 133 | + "$LOONGFORGE_PATH/loongforge/embodied/train.py" \ |
| 134 | + "${TRAINING_ARGS[@]}" \ |
| 135 | + "${ZERO1_ARGS[@]}" \ |
| 136 | + "${LR_ARGS[@]}" \ |
| 137 | + "${DATA_ARGS[@]}" \ |
| 138 | + "${LOGGING_ARGS[@]}" \ |
| 139 | + "${MODEL_DATA_OVERRIDES[@]+"${MODEL_DATA_OVERRIDES[@]}"}" \ |
| 140 | + "$@" 2>&1 | tee "$OUTPUT_DIR/$(basename "${BASH_SOURCE[0]}" .sh)_$(date +%Y%m%d_%H%M%S).log" |
0 commit comments