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