Skip to content

Commit 27cd61d

Browse files
committed
vla(fastwam): add ZeRO-1 SFT example
1 parent 01e25cb commit 27cd61d

2 files changed

Lines changed: 143 additions & 1 deletion

File tree

examples/embodied/fastwam/run_fastwam_sft_ddp.sh

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,10 @@ TOKENIZER_PATH=${TOKENIZER_PATH:-"/path/to/tokenizer"}
2121

2222
MASTER_PORT=${MASTER_PORT:-29519}
2323

24+
MODEL_NAME=${MODEL_NAME:-"fastwam"}
25+
2426
TRAINING_ARGS=(
25-
--model-name fastwam
27+
--model-name "$MODEL_NAME"
2628
--trainer-type FinetuneTrainer
2729
--distributed-strategy ddp
2830
--dtype bfloat16
Lines changed: 140 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,140 @@
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

Comments
 (0)