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:

uv pip install --torch-backend=auto "torch==2.13.0" → 2.13.0+cu129 (aarch64 wheel)

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):

cpunvidia grace, 64 × arm neoverse-v2 (aarch64), 1 thread/core, 1 socket
numalscpu 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 memory525 GiB lpddr5x
gpuNVIDIA 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 / cudadriver 570.148.08 (nvidia-smi reports cuda 12.8); toolkit nvcc 12.8.93 in /usr/local/cuda (Lambda Stack)
osubuntu 22.04.5, kernel 6.8.0-*-nvidia-64k: 64 KiB pages (getconf PAGESIZE = 65536), glibc 2.35
pythonsystem 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.12

2. 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-*-cu12 12.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" pyyaml

numpy 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_gemm
  • f1429a3 is the commit olmo-core’s dockerfile pins
  • --no-build-isolation builds against the torch you already installed
  • GROUPED_GEMM_CUTLASS=1 selects 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

componentstatusnotes
torch + cuda on aarch64working2.13.0+cu129, sm90, bf16 supported; bf16 8192³ matmul at 858 TFLOPS on an idle gpu (h100 pcie: 539)
torch.compile (inductor + triton)workinga compiled 60-layer model took 21–26 s for its first step from a cold cache
document-masked attentionworking (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 packagenot 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_mmworkingoutput identical to per-expert matmuls; backward works (docs)
moe: grouped_gemmworking with a workaroundsource build (step 4); output identical to per-expert matmuls
tokenizers / pyarrow / numpyworkingaarch64 wheels
memmap data loadingworking4 workers kept the gpu at 100% utilization
distributed checkpoints (dcp)workingno issues with 64 KiB pages
ncclworking (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.py

the 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$/hmicro-batch that fittok/s$ per 1B tokens
gh200 96 GB2.2933 (82.8 of 94.5 GiB reserved)159K in production, 169K in a sweep4.00
h100 pcie 80 GB3.292497.6K9.36
a100 40 GB (estimate)1.99n/a~64K8.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=auto and freeze from there.

  • uv venv over 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 at nice 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-64k kernel) 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_gemm in 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=zstd for 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

on this page