sskj-h3/throughput/sglang-profile/scripts/h3_profile/h3_flashattention_ncu_8gpu.sh
2026-08-31 15:57:13 +08:00

259 lines
8.9 KiB
Bash
Executable File

#!/usr/bin/env bash
# FlashAttention hardware-counter profile for MiniMax-H3 on 8x RTX 6000D.
# Four isolated TP2 services run concurrently:
# GPU 0-1: F3, cold-cache NCU replay
# GPU 2-3: F3, warm-cache NCU replay
# GPU 4-5: RVA, cold-cache NCU replay
# GPU 6-7: RVA, warm-cache NCU replay
set -Eeuo pipefail
ROOT=${ROOT:-/data/wxy}
RUN_ID=${RUN_ID:-h3-flashattention-ncu-$(date +%Y%m%d-%H%M%S)}
RESULT_ROOT=${RESULT_ROOT:-$ROOT/profile_results/$RUN_ID}
MODEL=${MODEL:-/data/hf_models/MiniMax-H3}
PYTHON=${PYTHON:-/root/.miniconda3/envs/sglang/bin/python}
SGLANG=${SGLANG:-/root/.miniconda3/envs/sglang/bin/sglang}
CLIENT=${CLIENT:-$ROOT/h3_profile/h3_profile_client.py}
NCU=${NCU:-/usr/local/cuda-13.2/bin/ncu}
MEDIA_BIN_DIR=${MEDIA_BIN_DIR:-/root/.miniconda3/envs/deploy/bin}
STEPS=${STEPS:-20}
DURATION=${DURATION:-5}
WARMUP_STEPS=${WARMUP_STEPS:-5}
BASE_PORT=${BASE_PORT:-30010}
PORT_STRIDE=${PORT_STRIDE:-10}
MASTER_PORT_BASE=${MASTER_PORT_BASE:-33000}
SCHEDULER_PORT_BASE=${SCHEDULER_PORT_BASE:-34000}
SERVER_START_TIMEOUT=${SERVER_START_TIMEOUT:-1800}
REFERENCE_IMAGE=${REFERENCE_IMAGE:-$ROOT/sskj-MiniMax-H3/assets/reference_images/landscape_mountain_lake.jpg}
REFERENCE_IMAGES_DIR=${REFERENCE_IMAGES_DIR:-$ROOT/h3_profile/assets/reference_images_5}
REFERENCE_VIDEO_1S=${REFERENCE_VIDEO_1S:-$MODEL/assets/r2va.mp4}
REFERENCE_VIDEO_5S=${REFERENCE_VIDEO_5S:-$MODEL/assets/ref2va.mp4}
REFERENCE_VIDEO_10S=${REFERENCE_VIDEO_10S:-$MODEL/assets/h3_direct_768p.mp4}
FLASH_REGEX=${FLASH_REGEX:-pytorch_flash::flash_fwd_kernel}
F3_LAUNCH_SKIP=${F3_LAUNCH_SKIP:-3734}
RVA_LAUNCH_SKIP=${RVA_LAUNCH_SKIP:-3921}
LAUNCH_COUNT=${LAUNCH_COUNT:-1}
declare -a SERVER_PIDS=() CLIENT_PIDS=()
log(){ printf '[%s] %s\n' "$(date '+%F %T')" "$*" | tee -a "$RESULT_ROOT/run.log"; }
die(){ log "ERROR: $*"; exit 1; }
port_open(){
"$PYTHON" - "$1" <<'PY'
import socket, sys
s = socket.socket()
s.settimeout(0.3)
try:
s.connect(("127.0.0.1", int(sys.argv[1])))
except OSError:
raise SystemExit(1)
raise SystemExit(0)
PY
}
wait_healthy(){
local port=$1 pid=$2 logfile=$3 deadline=$((SECONDS+SERVER_START_TIMEOUT))
while ((SECONDS < deadline)); do
if curl -fsS --max-time 2 "http://127.0.0.1:${port}/health" >/dev/null 2>&1; then return 0; fi
if ! kill -0 "$pid" 2>/dev/null; then tail -120 "$logfile"; return 1; fi
sleep 5
done
tail -120 "$logfile"
return 1
}
stop_servers(){
local p deadline alive
for p in "${SERVER_PIDS[@]:-}"; do
kill -INT -- "-$p" 2>/dev/null || kill -INT "$p" 2>/dev/null || true
done
deadline=$((SECONDS+180))
while ((SECONDS < deadline)); do
alive=0
for p in "${SERVER_PIDS[@]:-}"; do kill -0 "$p" 2>/dev/null && alive=1; done
((alive == 0)) && break
sleep 2
done
for p in "${SERVER_PIDS[@]:-}"; do
if kill -0 "$p" 2>/dev/null; then kill -TERM -- "-$p" 2>/dev/null || true; sleep 3; fi
if kill -0 "$p" 2>/dev/null; then kill -KILL -- "-$p" 2>/dev/null || true; fi
wait "$p" 2>/dev/null || true
done
SERVER_PIDS=()
}
cleanup(){
local rc=$? p
trap - EXIT INT TERM
for p in "${CLIENT_PIDS[@]:-}"; do kill -TERM "$p" 2>/dev/null || true; done
stop_servers
exit "$rc"
}
trap cleanup EXIT INT TERM
start_server(){
local replica=$1 variant=$2 gpus=$3 scenario=$4 cache_control=$5 launch_skip=$6
local port=$((BASE_PORT+replica*PORT_STRIDE))
local master=$((MASTER_PORT_BASE+replica*PORT_STRIDE))
local scheduler=$((SCHEDULER_PORT_BASE+replica*PORT_STRIDE))
local dir="$RESULT_ROOT/instances/replica_${replica}_${scenario}_${cache_control}"
local logfile="$dir/server.log"
mkdir -p "$dir/outputs" "$dir/ncu" "$dir/perf"
for p in "$port" "$((port+1))" "$master" "$scheduler"; do
port_open "$p" && die "port occupied: $p"
done
local -a sections=(
--set detailed
--section SchedulerStats
--section WarpStateStats
--section InstructionStats
)
local -a ncu_cmd=(
"$NCU"
--force-overwrite
--target-processes all
--kernel-name-base demangled
--kernel-name "regex:${FLASH_REGEX}"
--launch-skip "$launch_skip"
--launch-count "$LAUNCH_COUNT"
--replay-mode kernel
--cache-control "$cache_control"
--clock-control boost
"${sections[@]}"
--export "$dir/ncu/server"
)
local -a serve_cmd=(
"$SGLANG" serve
--model-path "$MODEL"
--model-variant "$variant"
--backend sglang
--performance-mode speed
--num-gpus 2
--tp-size 2
--ulysses-degree 1
--use-fsdp-inference false
--enable-torch-compile false
--batching-max-size 1
--batching-delay-ms 0
--enable-layerwise-nvtx-marker
--host 0.0.0.0
--port "$port"
--master-port "$master"
--scheduler-port "$scheduler"
--output-path "$dir/outputs"
--log-requests
--log-requests-level 2
)
printf '%q ' env CUDA_VISIBLE_DEVICES="$gpus" "${ncu_cmd[@]}" "${serve_cmd[@]}" > "$dir/command.txt"
printf '\n' >> "$dir/command.txt"
log "start replica=$replica scenario=$scenario GPUs=$gpus cache=$cache_control skip=$launch_skip port=$port"
CUDA_VISIBLE_DEVICES="$gpus" PATH="$MEDIA_BIN_DIR:$PATH" PYTHONUNBUFFERED=1 \
TOKENIZERS_PARALLELISM=false SGLANG_USE_RUNAI_MODEL_STREAMER=false \
SGLANG_DIFFUSION_STAGE_LOGGING=1 SGLANG_DIFFUSION_SYNC_STAGE_PROFILING=0 \
NCCL_DEBUG=WARN setsid "${ncu_cmd[@]}" "${serve_cmd[@]}" > "$logfile" 2>&1 &
SERVER_PIDS+=("$!")
}
run_client(){
local replica=$1 scenario=$2
local port=$((BASE_PORT+replica*PORT_STRIDE))
local cache_control=all
((replica % 2 == 1)) && cache_control=none
local dir="$RESULT_ROOT/instances/replica_${replica}_${scenario}_${cache_control}"
"$PYTHON" "$CLIENT" \
--deployment "FLASH_NCU_TP2_${cache_control}" \
--scenario "$scenario" \
--host 127.0.0.1 \
--port "$port" \
--replica-index "$replica" \
--model "$MODEL" \
--reference-image "$REFERENCE_IMAGE" \
--reference-images-dir "$REFERENCE_IMAGES_DIR" \
--reference-video-1s "$REFERENCE_VIDEO_1S" \
--reference-video-5s "$REFERENCE_VIDEO_5S" \
--reference-video-10s "$REFERENCE_VIDEO_10S" \
--steps "$STEPS" \
--duration "$DURATION" \
--short-edge 768 \
--aspect-ratio 16:9 \
--seed 1101 \
--repeats 1 \
--warmup 1 \
--warmup-steps "$WARMUP_STEPS" \
--perf-dir "$dir/perf" \
--output "$dir/results.jsonl" > "$dir/client.log" 2>&1
}
validate_report(){
local replica=$1 scenario=$2 expected_grid=$3
local cache_control=all
((replica % 2 == 1)) && cache_control=none
local dir="$RESULT_ROOT/instances/replica_${replica}_${scenario}_${cache_control}"
local report
report=$(find "$dir/ncu" -maxdepth 1 -name '*.ncu-rep' -print -quit)
[[ -n "$report" ]] || return 1
"$NCU" --import "$report" --page details --print-details all --csv > "$dir/ncu/details.csv" 2> "$dir/ncu/export.log"
grep -q "${expected_grid}, 1, 28" "$dir/ncu/details.csv" || grep -q "(${expected_grid}, 1, 28)" "$dir/ncu/details.csv"
}
mkdir -p "$RESULT_ROOT"/{metadata,instances,system}
cp "$0" "$RESULT_ROOT/metadata/$(basename "$0")"
date -Ins > "$RESULT_ROOT/metadata/date.txt"
nvidia-smi --query-gpu=index,name,uuid,memory.total,pci.bus_id,driver_version --format=csv > "$RESULT_ROOT/metadata/gpus.csv"
nvidia-smi topo -m > "$RESULT_ROOT/metadata/topology.txt"
"$NCU" --version > "$RESULT_ROOT/metadata/ncu-version.txt"
"$PYTHON" -m pip freeze > "$RESULT_ROOT/metadata/pip-freeze.txt"
log "launching four TP2 NCU services"
start_server 0 FL2VA 0,1 F3 all "$F3_LAUNCH_SKIP"
start_server 1 FL2VA 2,3 F3 none "$F3_LAUNCH_SKIP"
start_server 2 Ref2VA 4,5 RVA_EMBEDDED all "$RVA_LAUNCH_SKIP"
start_server 3 Ref2VA 6,7 RVA_EMBEDDED none "$RVA_LAUNCH_SKIP"
for replica in 0 1 2 3; do
port=$((BASE_PORT+replica*PORT_STRIDE))
scenario=F3
((replica >= 2)) && scenario=RVA_EMBEDDED
cache_control=all
((replica % 2 == 1)) && cache_control=none
logfile="$RESULT_ROOT/instances/replica_${replica}_${scenario}_${cache_control}/server.log"
wait_healthy "$port" "${SERVER_PIDS[$replica]}" "$logfile" || die "server failed: $logfile"
log "replica=$replica healthy port=$port"
done
nvidia-smi dmon -s pucvmet -d 2 -o DT > "$RESULT_ROOT/system/nvidia-dmon.log" 2>&1 &
DMON_PID=$!
log "starting F3/F3/RVA/RVA clients concurrently"
run_client 0 F3 & CLIENT_PIDS+=("$!")
run_client 1 F3 & CLIENT_PIDS+=("$!")
run_client 2 RVA_EMBEDDED & CLIENT_PIDS+=("$!")
run_client 3 RVA_EMBEDDED & CLIENT_PIDS+=("$!")
failed=0
for p in "${CLIENT_PIDS[@]}"; do wait "$p" || failed=1; done
CLIENT_PIDS=()
kill -TERM "$DMON_PID" 2>/dev/null || true
wait "$DMON_PID" 2>/dev/null || true
((failed == 0)) || die "one or more clients failed"
log "clients completed; stopping services so NCU can finalize reports"
stop_servers
validation_failed=0
validate_report 0 F3 327 || validation_failed=1
validate_report 1 F3 327 || validation_failed=1
validate_report 2 RVA_EMBEDDED 638 || validation_failed=1
validate_report 3 RVA_EMBEDDED 638 || validation_failed=1
((validation_failed == 0)) || die "NCU report missing or captured the wrong FlashAttention grid"
touch "$RESULT_ROOT/DONE"
log "DONE: $RESULT_ROOT"