all_code_base / lrm /flux /train_flux.sh
aryadomain's picture
Add files using upload-large-folder tool
533920b verified
Raw
History Blame Contribute Delete
9.32 kB
if [[ -t 1 && -n "${TERM:-}" ]]; then
clear
fi
set -euo pipefail
# Shared Hugging Face cache root used across nodes.
HF_HUB_CACHE_DIR="/scratch/rr81/ma5430/.cache/huggingface/hub"
export HF_HUB_CACHE="$HF_HUB_CACHE_DIR"
export HUGGINGFACE_HUB_CACHE="$HF_HUB_CACHE_DIR"
export HF_HOME="$(dirname "$HF_HUB_CACHE_DIR")"
export HF_DATASETS_OFFLINE=1
export HF_METRICS_OFFLINE=1
export HF_MODULES_OFFLINE=1
export TRANSFORMERS_OFFLINE=1
export DIFFUSERS_OFFLINE=1
export HF_HUB_OFFLINE=1
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
cd "$SCRIPT_DIR"
TMPDIR="${TMPDIR:-/tmp/flux_tmp}"
mkdir -p "$TMPDIR"
export TMPDIR
export TRITON_CACHE_DIR="$TMPDIR/triton_cache"
mkdir -p "$TRITON_CACHE_DIR"
CONFIG_PATH="$SCRIPT_DIR/trainer/conf"
# =========================
# Editable default values
# =========================
# Run profile selector.
# Options: main, quick
# - main: full/default multi-GPU training settings
# - quick: tiny smoke run settings (1 step, test_unique split)
RUN_PROFILE="main"
# Conda environment path that contains accelerate + training dependencies.
# Other examples: /opt/conda/envs/lrm, /home/user/miniconda3/envs/aev
CONDA_ENV_PATH="/g/data/rr81/aev"
# Main profile defaults.
MAIN_LAUNCH_MODE="distributed_sharded" # Options: distributed, distributed_sharded, data_parallel
MAIN_GPU_IDS="all" # Options: all, 0, 1, 0,1,2,3
MAIN_NUM_PROCESSES="1" # Typical values: 1, 2, 4, 8
MAIN_PROCESS_PORT="29100" # Change if port is in use
MAIN_PRECISION="BF16" # Options: BF16, FP16, FP32 (FP32 maps to NO mixed precision)
MAIN_ZERO3_OFFLOAD="cpu" # ZeRO-3 offload: none, cpu, nvme
MAIN_TRAIN_SPLIT="" # Example: train[:10%]
MAIN_VALID_SPLIT="validation_unique"
MAIN_TEST_SPLIT="test_unique"
# Tracking toggle for cluster/offline runs.
# 0: disable external trackers (recommended for unattended runs)
# 1: keep configured tracker (e.g., wandb, requires login/API key)
ENABLE_TRACKING="0"
# Quick profile defaults.
# data_parallel uses all visible GPUs for one process and is handy for quick bring-up checks.
# Note: DataParallel replicates the full model on each GPU (it does not shard model weights).
QUICK_LAUNCH_MODE="distributed_sharded" # Options: distributed, distributed_sharded, data_parallel
QUICK_GPU_IDS="all" # all, 0, 1, 0,1,2,3
QUICK_NUM_PROCESSES="8" # Used by distributed and distributed_sharded
QUICK_PROCESS_PORT="29100" # Change if port is in use
QUICK_PRECISION="FP16" # Options: BF16, FP16, FP32 (FP32 maps to NO mixed precision)
QUICK_ZERO3_OFFLOAD="cpu" # ZeRO-3 offload: none, cpu, nvme
QUICK_TRAIN_SPLIT="test_unique"
QUICK_VALID_SPLIT="test_unique"
QUICK_TEST_SPLIT="test_unique"
# Pseudo preference CSV path.
# Primary expected file: flux/vqa_aes_clip_score_mp.csv
# Alternate example: /efs/drsanny/data/vqa_aes_clip_score_mp.csv
PSEUDO_PREFERENCE_PATH="$SCRIPT_DIR/vqa_aes_clip_score_mp.csv"
# Fallback to parent path if needed.
if [[ ! -r "$PSEUDO_PREFERENCE_PATH" && -r "$SCRIPT_DIR/../vqa_aes_clip_score_mp.csv" ]]; then
PSEUDO_PREFERENCE_PATH="$SCRIPT_DIR/../vqa_aes_clip_score_mp.csv"
fi
case "$RUN_PROFILE" in
main)
LAUNCH_MODE="$MAIN_LAUNCH_MODE"
GPU_IDS="$MAIN_GPU_IDS"
NUM_PROCESSES="$MAIN_NUM_PROCESSES"
MAIN_PROCESS_PORT="$MAIN_PROCESS_PORT"
PRECISION="$MAIN_PRECISION"
ZERO3_OFFLOAD="$MAIN_ZERO3_OFFLOAD"
TRAIN_SPLIT="$MAIN_TRAIN_SPLIT"
VALID_SPLIT="$MAIN_VALID_SPLIT"
TEST_SPLIT="$MAIN_TEST_SPLIT"
IS_QUICK_RUN="0"
;;
quick)
LAUNCH_MODE="$QUICK_LAUNCH_MODE"
GPU_IDS="$QUICK_GPU_IDS"
NUM_PROCESSES="$QUICK_NUM_PROCESSES"
MAIN_PROCESS_PORT="$QUICK_PROCESS_PORT"
PRECISION="$QUICK_PRECISION"
ZERO3_OFFLOAD="$QUICK_ZERO3_OFFLOAD"
TRAIN_SPLIT="$QUICK_TRAIN_SPLIT"
VALID_SPLIT="$QUICK_VALID_SPLIT"
TEST_SPLIT="$QUICK_TEST_SPLIT"
IS_QUICK_RUN="1"
;;
*)
echo "[train_flux.sh] Invalid RUN_PROFILE='$RUN_PROFILE'. Use: main or quick." >&2
exit 1
;;
esac
case "${PRECISION^^}" in
BF16)
HYDRA_MIXED_PRECISION="BF16"
;;
FP16)
HYDRA_MIXED_PRECISION="FP16"
;;
FP32|NO|NONE)
HYDRA_MIXED_PRECISION="NO"
;;
*)
echo "[train_flux.sh] Invalid precision '$PRECISION'. Use: BF16, FP16, FP32." >&2
exit 1
;;
esac
ACCELERATE_MIXED_PRECISION="${HYDRA_MIXED_PRECISION,,}"
case "${LAUNCH_MODE}" in
distributed|distributed_sharded|data_parallel)
;;
*)
echo "[train_flux.sh] Invalid launch mode '$LAUNCH_MODE'. Use: distributed, distributed_sharded, or data_parallel." >&2
exit 1
;;
esac
if [[ "$LAUNCH_MODE" == "data_parallel" ]]; then
echo "[train_flux.sh] Launch mode=data_parallel on GPU_IDS=$GPU_IDS"
echo "[train_flux.sh] Note: DataParallel does not shard model weights. If a single GPU cannot hold the model, use distributed mode with sharding." >&2
fi
OVERRIDES=()
OVERRIDES+=("accelerator.mixed_precision=${HYDRA_MIXED_PRECISION}")
if [[ "$ENABLE_TRACKING" != "1" ]]; then
export WANDB_DISABLED=true
OVERRIDES+=("accelerator.log_with=null")
fi
if [[ "$LAUNCH_MODE" == "distributed_sharded" ]]; then
if [[ "$GPU_IDS" == "all" && "$NUM_PROCESSES" == "1" ]]; then
if command -v nvidia-smi >/dev/null 2>&1; then
NUM_PROCESSES="$(nvidia-smi -L | wc -l | tr -d ' ')"
echo "[train_flux.sh] distributed_sharded auto-set NUM_PROCESSES=$NUM_PROCESSES (all visible GPUs)"
fi
fi
case "${ZERO3_OFFLOAD,,}" in
none|cpu|nvme)
;;
*)
echo "[train_flux.sh] Invalid ZERO3 offload '$ZERO3_OFFLOAD'. Use: none, cpu, nvme." >&2
exit 1
;;
esac
echo "[train_flux.sh] Launch mode=distributed_sharded (DeepSpeed ZeRO-3)"
# Avoid MPI auto-discovery on single-node/single-process runs.
export MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
export MASTER_PORT="${MASTER_PORT:-$MAIN_PROCESS_PORT}"
export RANK="${RANK:-0}"
export LOCAL_RANK="${LOCAL_RANK:-0}"
export WORLD_SIZE="${WORLD_SIZE:-$NUM_PROCESSES}"
export DEEPSPEED_USE_MPI=0
OVERRIDES+=("accelerator=deepspeed")
OVERRIDES+=("criterion.is_distributed=true")
OVERRIDES+=("accelerator.deepspeed.zero_optimization.stage=3")
OVERRIDES+=("accelerator.deepspeed.gradient_accumulation_steps=1")
if [[ "${ZERO3_OFFLOAD,,}" != "none" ]]; then
OVERRIDES+=("+accelerator.deepspeed.zero_optimization.offload_param.device=${ZERO3_OFFLOAD}")
OVERRIDES+=("+accelerator.deepspeed.zero_optimization.offload_optimizer.device=${ZERO3_OFFLOAD}")
fi
fi
if [[ "$LAUNCH_MODE" == "distributed" ]]; then
OVERRIDES+=("accelerator=debug")
if [[ "$NUM_PROCESSES" == "1" ]]; then
OVERRIDES+=("criterion.is_distributed=false")
fi
fi
if [[ -n "$PSEUDO_PREFERENCE_PATH" && -r "$PSEUDO_PREFERENCE_PATH" ]]; then
OVERRIDES+=("dataset.pseudo_preference_path=${PSEUDO_PREFERENCE_PATH}")
elif [[ "$IS_QUICK_RUN" != "1" ]]; then
echo "[train_flux.sh] Warning: pseudo preference CSV not readable at '$PSEUDO_PREFERENCE_PATH'. Disabling pseudo preference filter for this run." >&2
OVERRIDES+=("dataset.keep_only_with_pesudo_preference=false")
fi
if [[ -n "$TRAIN_SPLIT" ]]; then
OVERRIDES+=("dataset.train_split_name=${TRAIN_SPLIT}")
fi
if [[ -n "$VALID_SPLIT" ]]; then
OVERRIDES+=("dataset.valid_split_name=${VALID_SPLIT}")
fi
if [[ -n "$TEST_SPLIT" ]]; then
OVERRIDES+=("dataset.test_split_name=${TEST_SPLIT}")
fi
if [[ "$IS_QUICK_RUN" == "1" ]]; then
OVERRIDES+=(
"dataset.train_split_name=test_unique"
"dataset.valid_split_name=test_unique"
"dataset.test_split_name=test_unique"
"dataset.keep_only_with_pesudo_preference=false"
"model.image_size=256"
"accelerator.max_steps=1"
"accelerator.validate_steps=1"
"accelerator.save_steps=0"
"dataset.batch_size=1"
)
fi
if [[ "$LAUNCH_MODE" == "distributed" || "$LAUNCH_MODE" == "distributed_sharded" ]]; then
conda run --no-capture-output -p "$CONDA_ENV_PATH" accelerate launch \
--dynamo_backend no \
--mixed_precision "$ACCELERATE_MIXED_PRECISION" \
--gpu_ids "$GPU_IDS" \
--num_processes "$NUM_PROCESSES" \
--num_machines 1 \
--main_process_port "$MAIN_PROCESS_PORT" \
trainer/scripts/train.py \
--config-path "$CONFIG_PATH" \
--config-name step_flux_base \
"${OVERRIDES[@]}"
else
OVERRIDES+=("accelerator=debug")
OVERRIDES+=("criterion.is_distributed=false")
if [[ "$GPU_IDS" == "all" ]]; then
conda run --no-capture-output -p "$CONDA_ENV_PATH" \
env USE_DATA_PARALLEL=1 python trainer/scripts/train.py \
--config-path "$CONFIG_PATH" \
--config-name step_flux_base \
"${OVERRIDES[@]}"
else
conda run --no-capture-output -p "$CONDA_ENV_PATH" \
env CUDA_VISIBLE_DEVICES="$GPU_IDS" USE_DATA_PARALLEL=1 python trainer/scripts/train.py \
--config-path "$CONFIG_PATH" \
--config-name step_flux_base \
"${OVERRIDES[@]}"
fi
fi