File size: 4,232 Bytes
0fe7113
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
# shellcheck shell=bash
# Shared environment for every Predictor shell script in this repository.
#
#   source "$(dirname "${BASH_SOURCE[0]}")/env.sh"
#
# Every value can be overridden from the caller's environment, e.g.
#   GPUS="0 1" DATA_ROOT=/some/disk bash scripts/run_stage1_offline_four_gpu_and_restore.sh

REPO="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"

# Python environment (torch 2.5.1 + flash-attn 2.8.3).
PYTHON_BIN="${PYTHON_BIN:-/local/zoubin/cz/envs/self_forcing/bin/python}"
# The env's torchrun wrapper has a stale shebang; launch the module directly.
TORCHRUN_CMD=("$PYTHON_BIN" -m torch.distributed.run)
VBENCH_PYTHON="${VBENCH_PYTHON:-/local/zoubin/cz/envs/vbench_eval/bin/python}"

# Sibling checkouts and model assets.  Inside the repo, checkpoints/chunkwise,
# wan_models and prompts/* are symlinks into ../Causal-Forcing and ../Self-Forcing.
# Self-Forcing-a ships scripts/evaluate_single_block_fppf.py (the FPPF reference
# evaluator); fall back to the older Self-Forcing checkout otherwise.
if [[ -z "${SELF_FORCING_ROOT:-}" ]]; then
    for candidate in "$(cd "$REPO/.." && pwd)/Self-Forcing-a" "$(cd "$REPO/.." && pwd)/Self-Forcing"; do
        if [[ -f "$candidate/scripts/evaluate_single_block_fppf.py" ]]; then
            SELF_FORCING_ROOT=$candidate; break
        fi
    done
    SELF_FORCING_ROOT="${SELF_FORCING_ROOT:-$(cd "$REPO/.." && pwd)/Self-Forcing}"
fi
# Predictor input fusion for Stage-1 launchers: self_forcing | disca | atc.
INPUT_VARIANT="${INPUT_VARIANT:-self_forcing}"
# Extra CLI flags forwarded to the Stage-1 trainers (e.g. "--atc_previous_scope last_frame").
STAGE1_EXTRA_ARGS="${STAGE1_EXTRA_ARGS:-}"
CAUSAL_CHECKPOINT="${CAUSAL_CHECKPOINT:-$REPO/checkpoints/chunkwise/causal_forcing.pt}"
CAUSAL_CONFIG="${CAUSAL_CONFIG:-$REPO/configs/causal_forcing_dmd_chunkwise.yaml}"
WAN_14B_DIR="${WAN_14B_DIR:-$REPO/wan_models/Wan2.1-T2V-14B}"
PROMPT_PATH="${PROMPT_PATH:-$REPO/prompts/vidprom_filtered_extended.txt}"
VALIDATION_PROMPT_PATH="${VALIDATION_PROMPT_PATH:-$REPO/prompts/MovieGenVideoBench_extended.txt}"
MOVIEBENCH_ORIGINAL="${MOVIEBENCH_ORIGINAL:-$REPO/prompts/MovieGenVideoBench.txt}"
MOVIEBENCH_EXTENDED="${MOVIEBENCH_EXTENDED:-$REPO/prompts/MovieGenVideoBench_extended.txt}"

# Where offline datasets and training/evaluation runs are written.
DATA_ROOT="${DATA_ROOT:-$REPO/data}"
OUTPUT_ROOT="${OUTPUT_ROOT:-$REPO/output}"

# Physical GPU indices, space separated.  GPUs 0-3 on this host are shared with
# other jobs, so the default uses the idle upper half.
GPUS="${GPUS:-4 5 6 7}"
read -r -a GPU_LIST <<< "$GPUS"
NUM_GPUS=${#GPU_LIST[@]}
GPUS_CSV=$(IFS=,; echo "${GPU_LIST[*]}")

export OMP_NUM_THREADS="${OMP_NUM_THREADS:-4}"
export MKL_NUM_THREADS="${MKL_NUM_THREADS:-4}"

# The login shell points HF_HOME/TORCH_HOME at an NVMe path that does not exist
# on this host; use a writable local cache instead (lpips/VBench weights).
CACHE_ROOT="${CACHE_ROOT:-/local/zoubin/cz/cache}"
export HF_HOME="${HF_HOME_OVERRIDE:-$CACHE_ROOT/huggingface}"
export HF_HUB_CACHE="$HF_HOME/hub"
export HUGGINGFACE_HUB_CACHE="$HF_HOME/hub"
export TORCH_HOME="$CACHE_ROOT/torch"
export SELF_FORCING_ROOT
mkdir -p "$DATA_ROOT" "$OUTPUT_ROOT" "$HF_HOME/hub" "$TORCH_HOME"

if [[ ! -x "$PYTHON_BIN" ]]; then
    echo "[env] python not found: $PYTHON_BIN (set PYTHON_BIN)" >&2
    return 1 2>/dev/null || exit 1
fi

# --- helpers ---------------------------------------------------------------

# log <master_log> <message...>
log() {
    local file=$1; shift
    printf '%s %s\n' "$(date --iso-8601=seconds)" "$*" | tee -a "$file"
}

# strided_ids <shard_index> <num_shards> <count>  -> "i i+n i+2n ..."
strided_ids() {
    local shard=$1 shards=$2 count=$3 ids=()
    for ((id=shard; id<count; id+=shards)); do ids+=("$id"); done
    echo "${ids[@]}"
}

# wait_all <master_log> <label> pid...  -> returns 1 if any child failed
wait_all() {
    local file=$1 label=$2; shift 2
    local status=0 index=0 pid
    for pid in "$@"; do
        if wait "$pid"; then
            log "$file" "[done] $label shard=$index"
        else
            log "$file" "[failed] $label shard=$index exit=$?"
            status=1
        fi
        index=$((index + 1))
    done
    return "$status"
}