gh200 setup for pytorch training (october 2026)
on this page
| tested on: | Lambda 1× GH200 96 GB, driver 570, Ubuntu 22.04, 64 KiB-page kernel (2026-10-04) |
| torch: |
|
| torch.compile: | works; triton 3.7.1 ships aarch64 wheels, no overrides needed |
| flash-attn: | not needed; use torch.nn.attention.varlen.varlen_attn for document masking |
| moe kernels: | F.grouped_mm works out of the box; grouped_gemm builds from source in ~1.5 min |
| main trap: | don’t reuse an x86 freeze; cu130 wheels need driver ≥ 580, and Lambda’s GH200 image ships 570 |
overview
this page covers a working pytorch training environment on a single nvidia gh200, set up in october 2026. it replaces the november 2025 gh200 page, which covered torch 2.9, cu128 and a triton override, none of which still apply.
practically speaking, the aarch64 story is now uneventful: every package below installed from a prebuilt wheel except one optional moe kernel. what remains is knowing which cuda build to pick, which packages you can skip, and a handful of quirks of the hardware.
each component below is marked working, working with a workaround or not working, based on runs on a real box rather than on release notes.
the box
measured on a Lambda gpu_1x_gh200 instance ($2.29/h at the time):
| cpu | nvidia grace, 64 × arm neoverse-v2 (aarch64), 1 thread/core, 1 socket |
| numa | lscpu reports 9 nodes, but all 64 cpus are on node 0; the rest are cpu-less nodes, the gpu’s hbm exposed to the os over nvlink-c2c |
| host memory | 525 GiB lpddr5x |
| gpu | NVIDIA GH200 480GB, sm_90, 97,871 MiB hbm3 (94.5 GiB usable by pytorch), 700 W cap, addressing mode ATS (cpu and gpu share one address space) |
| driver / cuda | driver 570.148.08 (nvidia-smi reports cuda 12.8); toolkit nvcc 12.8.93 in /usr/local/cuda (Lambda Stack) |
| os | ubuntu 22.04.5, kernel 6.8.0-*-nvidia-64k: 64 KiB pages (getconf PAGESIZE = 65536), glibc 2.35 |
| python | system 3.10, used only to bootstrap uv |
the “480GB” in the device name is the superchip total (cpu + gpu memory), not the gpu’s memory.
install
1. uv and python
python3 -m pip install --user uv # uv 0.12.x, aarch64-unknown-linux-gnu
export PATH=$HOME/.local/bin:/usr/local/cuda/bin:$PATH
uv venv --python 3.12 ~/venvs/train # uv downloads its own cpython 3.122. torch
PY=~/venvs/train/bin/python
uv pip install --python "$PY" --torch-backend=auto "torch==2.13.0"--torch-backend=auto reads the driver version, and driver 570 (cuda 12.8) leads it to the cu129 index through cuda minor-version compatibility. on aarch64 that resolves torch-2.13.0+cu129, a manylinux_2_28_aarch64 wheel from download.pytorch.org/whl/cu129, along with:
nvidia-*-cu1212.9 runtime wheels (cuBLAS 12.9.1.4, cuDNN 9.20.0.48, NCCL 2.29.7)triton 3.7.1(aarch64)
no pyproject.toml index overrides and no triton workaround are needed. if you pin an index instead of using auto, use cu129 on a 570 driver. cu130 (the default build since torch 2.11, see pytorch setup with uv) needs driver ≥ 580.
3. the rest of a training stack
uv pip install --python "$PY" --torch-backend=auto numpy tokenizers pyarrow "transformers==5.17.0" pyyamlnumpy 2.5.3, tokenizers 0.23.2, pyarrow 25.0.1 and transformers all come from prebuilt aarch64 wheels, so nothing compiles. pass --torch-backend=auto on every install that might pull torch in, so a dependency can’t swap in a different build.
an evaluation venv with lm-eval==0.4.13 needs transformers 4.x (4.57.6 worked), which in turn pins tokenizers 0.22.x and huggingface-hub 0.36.x. keep it in a separate venv rather than downgrading the training one.
4. optional: grouped_gemm for moe
skip this step unless your moe code needs grouped_gemm specifically, e.g. olmo-core’s DroplessMoE. pytorch’s own F.grouped_mm covers the same operation (see below).
uv pip install --python "$PY" setuptools wheel ninja
git clone https://github.com/tgale96/grouped_gemm.git ~/src/grouped_gemm
cd ~/src/grouped_gemm && git checkout f1429a3 && git submodule update --init --recursive
GROUPED_GEMM_CUTLASS=1 TORCH_CUDA_ARCH_LIST=9.0 MAX_JOBS=24 \
uv pip install --python "$PY" --no-build-isolation --no-deps ~/src/grouped_gemmf1429a3is the commit olmo-core’s dockerfile pins--no-build-isolationbuilds against the torch you already installedGROUPED_GEMM_CUTLASS=1selects the cutlass path; without it, the kernel rejects batch-size tensors that live on the gpu- with nvcc 12.8, the build took about 1 min 18 s using 24 grace cores
5. optional: xgrammar
uv pip install xgrammar installs 0.2.8 from a prebuilt aarch64 wheel.
component matrix
| component | status | notes |
|---|---|---|
| torch + cuda on aarch64 | working | 2.13.0+cu129, sm90, bf16 supported; bf16 8192³ matmul at 858 TFLOPS on an idle gpu (h100 pcie: 539) |
torch.compile (inductor + triton) | working | a compiled 60-layer model took 21–26 s for its first step from a cold cache |
| document-masked attention | working (pytorch-native) | torch.nn.attention.varlen.varlen_attn, which is in-tree and derived from fa2; bf16 gradients match per-document sdpa to about 3e-3 relative error, the same as on x86 h100/a100 |
flash-attn package | not installed (not needed) | no aarch64 wheels upstream, and source builds on gh200 are a known pain (#1866, #2036, #2045); fa3 is the next step if you need more attention speed |
moe: F.grouped_mm | working | output identical to per-expert matmuls; backward works (docs) |
moe: grouped_gemm | working with a workaround | source build (step 4); output identical to per-expert matmuls |
| tokenizers / pyarrow / numpy | working | aarch64 wheels |
| memmap data loading | working | 4 workers kept the gpu at 100% utilization |
| distributed checkpoints (dcp) | working | no issues with 64 KiB pages |
| nccl | working (single gpu) | 2.29.7 initializes the 1-rank process group that torchrun creates; multi-node not tested |
verify
check_gh200.py: run it from outside any source tree. it prints a pass/fail line per check and exits non-zero if any check fails:
"""gh200 env checks: device, bf16 matmul, torch.compile, varlen attention, grouped gemm."""
import sys, time, torch, torch.nn.functional as F
ok = True
def rep(name, good, msg=""):
global ok
ok &= bool(good)
print(f"{'PASS' if good else 'FAIL'} {name} {msg}", flush=True)
if not torch.cuda.is_available():
rep("device", False, f"torch {torch.__version__}: no cuda")
sys.exit(1)
p = torch.cuda.get_device_properties(0)
rep("device", True, f"{p.name} sm{p.major}{p.minor} {p.total_memory / 2**30:.1f} GiB "
f"torch {torch.__version__} cuda {torch.version.cuda}")
rep("bf16", torch.cuda.is_bf16_supported())
# bf16 matmul throughput
a = torch.randn(8192, 8192, device="cuda", dtype=torch.bfloat16); b = torch.randn_like(a)
for _ in range(3): a @ b
torch.cuda.synchronize(); t = time.time()
for _ in range(20): a @ b
torch.cuda.synchronize()
rep("bf16 matmul", True, f"{2 * 8192**3 * 20 / (time.time() - t) / 1e12:.0f} TFLOPS")
# torch.compile (inductor + triton on aarch64)
f = torch.compile(lambda x: F.gelu(x) * 2 + x.sin())
x = torch.randn(4096, 4096, device="cuda")
rep("torch.compile", torch.allclose(f(x), F.gelu(x) * 2 + x.sin(), atol=1e-5))
# varlen attention (document masking) vs per-document sdpa
from torch.nn.attention.varlen import varlen_attn
q = torch.randn(4096, 9, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True)
cu = torch.tensor([0, 1000, 4096], device="cuda", dtype=torch.int32)
o = varlen_attn(q, q, q, cu, cu, 4096, 4096, window_size=(-1, 0))
def doc(s, e):
t = q[s:e].transpose(0, 1)[None]
return F.scaled_dot_product_attention(t, t, t, is_causal=True)[0].transpose(0, 1)
ref = torch.cat([doc(0, 1000), doc(1000, 4096)])
o.float().pow(2).mean().backward()
diff = (o - ref).abs().max().item()
rep("varlen_attn", diff < 2e-2, f"max|diff| vs per-doc sdpa {diff:.2e}")
# moe grouped gemm vs per-expert matmuls
E, d, h = 6, 576, 1536
sizes = torch.tensor([3000, 0, 5000, 1234, 7000, 150])
offs = torch.cumsum(sizes, 0).cuda().to(torch.int32)
xx = torch.randn(int(sizes.sum()), d, device="cuda", dtype=torch.bfloat16, requires_grad=True)
w = torch.randn(E, h, d, device="cuda", dtype=torch.bfloat16, requires_grad=True)
r = torch.cat([c @ w[i].t() for i, c in enumerate(torch.split(xx, sizes.tolist()))])
y = F.grouped_mm(xx, w.transpose(-2, -1), offs=offs)
y.float().pow(2).mean().backward()
rep("F.grouped_mm", (y - r).abs().max().item() == 0)
try:
import grouped_gemm
y2 = grouped_gemm.ops.gmm(xx, w, sizes.cuda(), trans_b=True)
y2.float().pow(2).mean().backward()
rep("grouped_gemm", (y2 - r).abs().max().item() == 0, grouped_gemm.__file__)
except ImportError:
print("SKIP grouped_gemm (not installed)")
sys.exit(0 if ok else 1)cd /tmp && ~/venvs/train/bin/python check_gh200.pythe backward passes use .pow(2).mean() rather than .sum() on purpose: .sum().backward() produces a broadcast (stride-0) gradient, which F.grouped_mm rejects. that only matters in micro-benchmarks; a real loss doesn’t produce that gradient.
throughput and cost
one data point: a 60-layer, ~1.76B-parameter dense decoder (d576, swiglu, seq 2048, document masking, compile, adamw, global batch 264 × 2048), run on one gpu.
| box | $/h | micro-batch that fit | tok/s | $ per 1B tokens |
|---|---|---|---|---|
| gh200 96 GB | 2.29 | 33 (82.8 of 94.5 GiB reserved) | 159K in production, 169K in a sweep | 4.00 |
| h100 pcie 80 GB | 3.29 | 24 | 97.6K | 9.36 |
| a100 40 GB (estimate) | 1.99 | n/a | ~64K | 8.64 |
on this workload, the gh200 cost less than half as much per token as either alternative. the extra 16 GB also allowed a larger micro-batch (33 vs 24). the gap between the production rate and the sweep rate comes from a mix of shorter documents and the single-core kernel-launch bottleneck described below.
gotchas
things that bit
don’t reuse an x86 freeze. a lockfile from an x86 workstation on cu130 won’t run on driver 570. resolve on the box with
--torch-backend=autoand freeze from there.uv venvover an existing venv recreates it. if you script the setup, guard it with[ -x "$VENV/bin/python" ] || uv venv ..., or install into a fresh path and compare freezes.the python main process is the bottleneck, not the cpu count. a single training process used about 94% of one grace core on kernel launches while the gpu sat at 100%. pinning training to 6 cores (
taskset -c 0-5: main process + 4 data-loader workers) was enough, which left the other ~58 cores free for builds and evals atnice 10.numa looks strange but is harmless. all cores are on node 0, and the hbm shows up as cpu-less numa nodes. there is nothing to tune for one gpu.
64 KiB pages (the
-nvidia-64kkernel) caused no problems with numpy memmaps or distributed checkpoints.importing a package from inside its source dir shadows the installed copy. this happens with
grouped_gemmin particular, so run checks from/tmp.the gpu is shared with side jobs. at 88% reserved, about 12 GB is left for evals or smoke tests next to a training run. plan side jobs to fit, or wait for the run to finish.
uploads are limited by your own uplink. use
rsync --compress-choice=zstdfor token data (it compressed to about 65% of raw), but not for checkpoints (about 92%, not worth the cpu). do data selection before upload, not on the box.
setup script
setup_gh200.sh puts the steps above in one script. it is safe to rerun: it won’t recreate a venv that already exists, and it skips grouped_gemm if that already imports.
#!/usr/bin/env bash
# usage: bash setup_gh200.sh [VENV=~/venvs/train] (GROUPED_GEMM=0 to skip the moe kernel)
set -euo pipefail
VENV=${1:-$HOME/venvs/train}; SRC=${SRC:-$HOME/src}
export PATH=$HOME/.local/bin:/usr/local/cuda/bin:$PATH CUDA_HOME=${CUDA_HOME:-/usr/local/cuda}
command -v uv >/dev/null || python3 -m pip install --user -q uv
[ -x "$VENV/bin/python" ] || uv venv -q --python 3.12 "$VENV"
PY=$VENV/bin/python
uv pip install -q --python "$PY" --torch-backend=auto "torch==2.13.0"
uv pip install -q --python "$PY" --torch-backend=auto numpy tokenizers pyarrow "transformers==5.17.0" pyyaml
if [ "${GROUPED_GEMM:-1}" = 1 ] && ! (cd /tmp && "$PY" -c "import grouped_gemm" 2>/dev/null); then
uv pip install -q --python "$PY" setuptools wheel ninja
[ -d "$SRC/grouped_gemm" ] || git clone -q https://github.com/tgale96/grouped_gemm.git "$SRC/grouped_gemm"
(cd "$SRC/grouped_gemm" && git checkout -q f1429a3 && git submodule update -q --init --recursive)
GROUPED_GEMM_CUTLASS=1 TORCH_CUDA_ARCH_LIST=9.0 MAX_JOBS=$(( $(nproc) / 2 )) \
uv pip install --python "$PY" --no-build-isolation --no-deps "$SRC/grouped_gemm"
fi
uv pip freeze --python "$PY" > "$VENV.freeze.txt"
echo "done: $VENV ($(wc -l < "$VENV.freeze.txt") packages)"the run this page is based on reran the script into a second venv path and got an identical uv pip freeze (apart from editable and local-path lines).
references
- pytorch setup with uv: index urls and backends for torch 2.13
- uv pytorch integration: how
--torch-backendresolves - pytorch varlen attention tutorial
torch.nn.functional.grouped_mm- tgale96/grouped_gemm
- nvidia grace hopper superchip
- pytorch on gh200 (november 2025): the earlier, now-outdated version of this page