#!/bin/bash # WorldMem diffusion multi-node training, launched MANUALLY on each node. # # Run this on every node, passing that node's rank as the only argument. # Node 0 must be the machine whose IP matches MASTER_ADDR. # # Usage (run on each node): # bash scripts/train_multinode.sh [exp_name] # # Example with 2 nodes: # # on master node: # bash scripts/train_multinode.sh 0 decmem # # on worker node: # bash scripts/train_multinode.sh 1 decmem # # Env overrides: # EXP_NAME experiment name / config stem (default: decmem) # NPROC_PER_NODE GPUs per node (default 8) # MASTER_PORT torchrun rendezvous port (default 29500) # RESUME_CKPT path to a checkpoint to fully resume from (optional) # ─── Topology (EDIT THESE) ─────────────────────────────────────────────────── NNODES=2 # TODO: total number of nodes MASTER_ADDR="" # TODO: IP of node_rank=0 MASTER_PORT=${MASTER_PORT:-29500} NPROC_PER_NODE=${NPROC_PER_NODE:-8} # ─── Node rank / experiment from CLI ───────────────────────────────────────── NODE_RANK=$1 EXP_NAME=${2:-${EXP_NAME:-decmem}} if [ -z "$NODE_RANK" ]; then echo "Error: NODE_RANK is required. Usage: bash $0 [exp_name]" exit 1 fi # ─── Paths / config ────────────────────────────────────────────────────────── CONFIG=configs/${EXP_NAME}.yaml LOGDIR=exps/${EXP_NAME} WANDB_SAVE_DIR=${LOGDIR}/wandb TORCHRUN=torchrun RESUME_CKPT=${RESUME_CKPT:-} mkdir -p "${LOGDIR}" # Optional resume: `RESUME_CKPT=/path/to/model.pt bash ` # Must be set on EVERY node (this script doesn't propagate env across nodes). RESUME_ARGS="" if [ -n "${RESUME_CKPT}" ]; then RESUME_ARGS="--resume-ckpt ${RESUME_CKPT}" echo "[resume] RESUME_CKPT = ${RESUME_CKPT}" fi # ─── NCCL (match train_worldmem_multinode_torchrun.sh) ─────────────────────── export NCCL_IB_DISABLE=0 export NCCL_IB_GID_INDEX=3 export NCCL_MIN_NCHANNELS=16 export NCCL_IB_HCA=mlx5 export NCCL_IB_QPS_PER_CONNECTION=4 export NCCL_IB_TIMEOUT=22 export NCCL_DEBUG=WARN echo "EXP_NAME = ${EXP_NAME}" echo "CONFIG = ${CONFIG}" echo "LOGDIR = ${LOGDIR}" echo "NNODES = ${NNODES}" echo "NPROC_PER_NODE = ${NPROC_PER_NODE}" echo "WORLD_SIZE = $((NNODES * NPROC_PER_NODE))" echo "MASTER_ADDR = ${MASTER_ADDR}" echo "MASTER_PORT = ${MASTER_PORT}" echo "RESUME_CKPT = ${RESUME_CKPT:-}" echo "Starting node with NODE_RANK=${NODE_RANK}" # ─── Launch torchrun ───────────────────────────────────────────────────────── ${TORCHRUN} \ --nnodes=${NNODES} \ --nproc_per_node=${NPROC_PER_NODE} \ --node_rank=${NODE_RANK} \ --master_addr=${MASTER_ADDR} \ --master_port=${MASTER_PORT} \ train.py \ --config_path ${CONFIG} \ --logdir ${LOGDIR} \ --wandb-save-dir ${WANDB_SAVE_DIR} \ ${RESUME_ARGS} \ > "${LOGDIR}/node_${NODE_RANK}_$(date +%Y%m%d_%H%M%S).log" 2>&1