worldmem-baseline-evals / DecMem /scripts /train_multinode.sh
BonanDing's picture
Add isolated Minecraft and RE10K baseline evaluation suite
59630ba verified
Raw History Blame Contribute Delete
3.43 kB
#!/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 <node_rank> [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 <node_rank> [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 <this_script> <node_rank>`
# 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:-<none>}"
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