-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathlaunch_ddp_on_aws.sh
More file actions
139 lines (108 loc) · 3.88 KB
/
Copy pathlaunch_ddp_on_aws.sh
File metadata and controls
139 lines (108 loc) · 3.88 KB
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
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
#!/usr/bin/env bash
#
# launch_ddp_on_aws.sh
# Discover EC2 instances by tag, assign NODE_RANK / MASTER_ADDR,
# and launch torchrun on each node via SSH.
#
# Usage:
# export CLUSTER_TAG_VALUE=mycluster1
# ./launch_ddp_on_aws.sh train.py --arg1 ... --argN ...
#
# Requirements:
# - All nodes have the same tag key/value (default key: TrainingCluster).
# - AWS CLI and jq installed on this node.
# - Instance profile (IAM role) allows ec2:DescribeInstances.
# - SSH key can be used to login as ec2-user to each node.
#
set -euo pipefail
########################
# Configurable parameters
########################
# Tag key / value to identify the training cluster
CLUSTER_TAG_KEY="${CLUSTER_TAG_KEY:-TrainingCluster}"
CLUSTER_TAG_VALUE="${CLUSTER_TAG_VALUE:-}"
# SSH settings
SSH_KEY_PATH="${SSH_KEY_PATH:-$HOME/.ssh/id_rsa}"
SSH_USER="${SSH_USER:-ec2-user}"
# torchrun / DDP settings
MASTER_PORT="${MASTER_PORT:-29500}"
GPUS_PER_NODE="${GPUS_PER_NODE:-}" # if empty, will auto-detect via nvidia-smi
########################
# Basic checks
########################
if [[ -z "$CLUSTER_TAG_VALUE" ]]; then
echo "ERROR: CLUSTER_TAG_VALUE is not set. Example:" >&2
echo " export CLUSTER_TAG_VALUE=mycluster1" >&2
exit 1
fi
if ! command -v aws &>/dev/null; then
echo "ERROR: aws CLI not found. Install awscli first." >&2
exit 1
fi
if ! command -v jq &>/dev/null; then
echo "ERROR: jq not found. Install jq first." >&2
exit 1
fi
if [[ "$#" -lt 1 ]]; then
echo "Usage: CLUSTER_TAG_VALUE=<value> $0 <train_script.py> [args...]" >&2
exit 1
fi
TRAIN_SCRIPT="$1"
shift
TRAIN_ARGS="$*"
########################
# Detect region from EC2 metadata
########################
echo "[launch_ddp_on_aws] Detecting region from instance metadata..."
REGION="$(curl -s http://169.254.169.254/latest/dynamic/instance-identity/document | jq -r .region)"
echo "[launch_ddp_on_aws] Using region: ${REGION}"
########################
# Discover instances by tag
########################
echo "[launch_ddp_on_aws] Discovering instances with tag ${CLUSTER_TAG_KEY}=${CLUSTER_TAG_VALUE}..."
HOSTS_JSON="$(
aws ec2 describe-instances \
--region "$REGION" \
--filters "Name=tag:${CLUSTER_TAG_KEY},Values=${CLUSTER_TAG_VALUE}" "Name=instance-state-name,Values=running" \
--query 'Reservations[].Instances[].PrivateIpAddress' \
--output json
)"
# Convert to sorted Bash array
mapfile -t HOSTS < <(echo "$HOSTS_JSON" | jq -r '.[]' | sort)
WORLD_SIZE="${#HOSTS[@]}"
if [[ "$WORLD_SIZE" -eq 0 ]]; then
echo "ERROR: No running instances found for tag ${CLUSTER_TAG_KEY}=${CLUSTER_TAG_VALUE}" >&2
exit 1
fi
echo "[launch_ddp_on_aws] Found ${WORLD_SIZE} hosts:"
printf ' %s\n' "${HOSTS[@]}"
MASTER_ADDR="${HOSTS[0]}"
echo "[launch_ddp_on_aws] MASTER_ADDR=${MASTER_ADDR}, MASTER_PORT=${MASTER_PORT}"
########################
# Determine GPUs per node
########################
if [[ -z "${GPUS_PER_NODE}" ]]; then
echo "[launch_ddp_on_aws] GPUS_PER_NODE is not set; auto-detecting on this node via nvidia-smi..."
if command -v nvidia-smi &>/dev/null; then
GPUS_PER_NODE="$(nvidia-smi -L | wc -l)"
else
echo "ERROR: nvidia-smi not found and GPUS_PER_NODE not set." >&2
exit 1
fi
fi
echo "[launch_ddp_on_aws] GPUS_PER_NODE=${GPUS_PER_NODE}"
########################
# Launch on each host via SSH
########################
echo "[launch_ddp_on_aws] Launching training on all hosts..."
for i in "${!HOSTS[@]}"; do
host="${HOSTS[$i]}"
rank="$i"
echo "[launch_ddp_on_aws] Launching NODE_RANK=${rank} on host ${host}..."
ssh -i "$SSH_KEY_PATH" -o StrictHostKeyChecking=no "${SSH_USER}@${host}" \
"export NNODES=${WORLD_SIZE} NODE_RANK=${rank} MASTER_ADDR=${MASTER_ADDR} MASTER_PORT=${MASTER_PORT} GPUS_PER_NODE=${GPUS_PER_NODE}; \
torchrun_ddp.sh ${TRAIN_SCRIPT} ${TRAIN_ARGS}" &
done
echo "[launch_ddp_on_aws] Waiting for all ranks to finish..."
wait
echo "[launch_ddp_on_aws] All ranks finished."