The AMD Instinct MI300X is not one die. It is eight compute dies, called XCDs, stacked on four I/O dies, with 192 GB of HBM3 around them. For the CUDA-shaped mental model of one GPU with one L2, that matters in three ways: the L2 is split into eight private 4 MB slices, a 256 MB Infinity Cache sits behind them, and by default workgroups are dealt to the XCDs round-robin. A GEMM tuned for a monolithic L2 can fetch the same operand tile eight times.
This article explains the package as kernels and serving stacks see it: the memory hierarchy, round-robin dispatch and how to remap around it, wave64 and the matrix instructions, a decode and prefill roofline, partition modes, the 8-GPU platform and the settings that decide your numbers. For porting (HIP, hipify, RCCL) see AMD Instinct and ROCm; for the memory-upgraded successor see MI325X.
The part in software terms
| Property | MI300X | What it means for software |
|---|---|---|
| Compute dies | 8 XCDs, 304 active CUs (38 per XCD) | Grids need well over 304 workgroups to fill the device |
| L2 cache | 4 MB per XCD, private to that XCD | Reuse only happens between workgroups on the same XCD |
| Infinity Cache | 256 MB, shared | Catches L2 misses before HBM; large working sets still benefit |
| HBM3 | 192 GB, 5.3 TB/s peak | A 70B model in BF16 fits on one GPU |
| LDS | 64 KB per CU | Bounds tile size and pipeline depth in shared-memory kernels |
| Peak dense FP16/BF16 | 1307.4 TFLOPS | Matrix-core figure, not vector ALU |
| Peak dense FP8 | 2614.9 TFLOPS | Uses the FNUZ FP8 variants, not OCP FP8 |
| Board power | 750 W | Sustained clocks depend on cooling and power caps |
| Peer links | 128 GB/s per GPU pair, 896 GB/s aggregate in an 8-GPU platform | Tensor parallelism inside one node |
These are AMD data sheet and ROCm documentation figures, quoted without sparsity. Peak numbers are ceilings: production GEMMs reach a fraction of them, and the fraction depends heavily on the tuning covered below. The gfx target for MI300X in ROCm is gfx942; check it with rocminfo or, in PyTorch, torch.cuda.get_device_properties(0).gcnArchName.
The memory hierarchy, level by level
Walk a load from the inside out. A compute unit has its registers and 64 KB of local data share (LDS), the equivalent of CUDA shared memory, managed explicitly by the kernel. Below that is the XCD's L2, 4 MB shared by the 38 CUs on that die and nothing else; the ROCm documentation describes it as coalescing all memory traffic for the die. A miss crosses the Infinity Fabric on the I/O dies to the 256 MB Infinity Cache, and only a miss there reaches HBM.
So the effective L2 for any workgroup is 4 MB, not 32. Raster orders that rely on a neighbouring workgroup having just fetched a panel only work if that neighbour ran on the same XCD. The Infinity Cache softens the penalty, since data that misses L2 may still avoid HBM, but the fabric hop still costs latency and bandwidth. The HBM architecture article covers the DRAM side.
Round-robin dispatch and tile reuse
In the default single-partition (SPX) mode the eight XCDs appear as one device, and AMD documents that workgroups launched to it are distributed round-robin across the XCDs: workgroup 0 to XCD 0, workgroup 7 to XCD 7, workgroup 8 back to XCD 0.
Consider a tiled GEMM whose program ids map row-major onto output tiles. Tiles 0 to 7 share one panel of A. On a monolithic GPU they share it in L2; on MI300X each lands on a different XCD, so the panel is fetched into eight L2 slices. Grouped tile orders are defeated for the same reason.
The fix is to change which logical tile each workgroup computes: if workgroup p runs on XCD p mod 8, give XCD x a contiguous range of tiles. This is arithmetic derived from the documented dispatch order, not a vendor API, and it composes with grouped ordering applied afterwards:
NUM_XCDS = 8
def remap_pid(pid, num_pids):
# pid runs on XCD pid % 8. Hand each XCD a contiguous block of logical tiles.
per_xcd = num_pids // NUM_XCDS
if pid >= per_xcd * NUM_XCDS: # leftover tail: leave unchanged
return pid
return (pid % NUM_XCDS) * per_xcd + pid // NUM_XCDS
# In a Triton kernel the same expression is applied to tl.program_id(0)
# before the usual grouped (GROUP_M) tile ordering.Checked over 64 program ids, the mapping is a permutation, with XCD 0 computing logical tiles 0 to 7 and XCD 1 tiles 8 to 15. It assumes SPX; in CPX each XCD is its own device and no remap is needed. Prefer the vendor GEMM where it covers your shape; the remap matters most for your own Triton and HIP kernels.
Worked example: L2 traffic for one tile row
Quantify it for one tile row of a BF16 GEMM with 256 x 256 output tiles and a K step of 64. Each K step reads a 256 x 64 slice of A and a 64 x 256 slice of B, 32 KB each. Take eight neighbouring tiles in one tile row: they share one A slice and need eight different B slices.
| Placement | L2 fills per K step | Data moved into L2 |
|---|---|---|
| Round-robin, 1 tile per XCD | 8 XCDs x (1 A + 1 B) | 16 slices, 512 KB |
| Remapped, 1 x 8 tiles on one XCD | 1 A + 8 B | 9 slices, 288 KB |
| Remapped, 2 x 4 block on one XCD | 2 A + 4 B | 6 slices, 192 KB |
That is 1.8x to 2.7x less traffic into L2 for the same work, all of it traffic that would otherwise be served by Infinity Cache or HBM. In a real kernel each XCD runs dozens of tiles at once, which is exactly why the contiguous block assignment pays off: the more tiles co-resident on one die, the more operand slices they share. Whether that becomes a measurable speed-up depends on whether the kernel was bandwidth- or latency-bound at the fabric, so measure before and after with the ROCm profiler rather than assuming the ratio.
Wave64 and the matrix instructions
CDNA executes wavefronts of 64 lanes, not warps of 32. Code ported from CUDA that hard-codes 32, in warp shuffles, ballot masks held in 32-bit integers, or reductions that stop at 16 lanes, produces wrong answers rather than slow ones. Query the width at run time (warpSize in HIP device code, hipDeviceProp_t::warpSize on the host) and use 64-bit masks.
Matrix work runs on MFMA instructions. AMD's MI300X tuning guide gives several Triton rules of thumb worth starting from: matrix_instr_nonkdim=16 (the 16 x 16 MFMA) typically outperforms 32 x 32 even for large tiles; num_stages=2 for single-GEMM kernels and 1 for kernels fusing two GEMMs, because the 64 KB LDS fills quickly; and waves_per_eu as a hint that asks the compiler to cut register use to reach a target occupancy.
@triton.autotune(configs=[
triton.Config({"BLOCK_M": 256, "BLOCK_N": 256, "BLOCK_K": 64, "GROUP_M": 8,
"matrix_instr_nonkdim": 16, "waves_per_eu": 2},
num_warps=8, num_stages=2),
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64, "GROUP_M": 8,
"matrix_instr_nonkdim": 16, "waves_per_eu": 0},
num_warps=4, num_stages=2),
], key=["M", "N", "K"])
@triton.jit
def matmul_kernel(...):
pid = tl.program_id(0)
pid = remap_pid_for_xcds(pid, num_pids) # the arithmetic shown above
...
A roofline for decode and prefill
Divide peak dense BF16 compute by HBM bandwidth: 1307.4 TFLOPS / 5.3 TB/s is about 247 FLOPs per byte. A kernel with lower arithmetic intensity is bandwidth-bound on HBM. Decoding one token with a dense transformer reads every weight once and does about two FLOPs per weight, so at batch size 1 the intensity is roughly 1 FLOP per byte in BF16, far below the ridge.
Worked out for a 70B-parameter model in BF16: 140 GB of weights must stream from HBM per decode step. At 5.3 TB/s that is at least 26 ms, so a single sequence cannot exceed about 38 tokens per second however good the kernels are, before counting KV-cache reads. Batching raises intensity because one weight read serves many sequences; capacity is what lets you batch, and 192 GB leaves roughly 50 GB, less runtime overhead, for KV cache after the weights on one GPU. Prefill, with thousands of tokens per weight read, sits above the ridge and is compute-bound, which is where MFMA tuning and FP8 pay off. The same analysis for Hopper is in the H200 article.
Partition modes in brief
The XCDs can also be exposed separately. In CPX mode each XCD appears as its own GPU, so one MI300X becomes eight devices. Memory has its own setting: NPS1 makes all HBM visible to every XCD, while NPS4 gives each quarter of the memory to the logical devices in its quadrant, and is only valid with CPX. AMD measured about 4,210 GB/s total stream bandwidth in CPX with NPS4 against about 4,010 GB/s with NPS1.
# Inspect, then change partitioning (requires root; drains the GPUs).
amd-smi static --partition
amd-smi set --gpu all --compute-partition CPX
amd-smi set --gpu all --memory-partition NPS4CPX suits many small independent jobs, such as small-model replicas, that cannot fill 304 CUs, and is wrong for one large model. AMD lists other intermediate modes; check your ROCm version's documentation before using them.
The 8-GPU platform
In the standard 8-GPU platform every GPU has a direct Infinity Fabric link to each of the other seven, at 128 GB/s per pair and 896 GB/s aggregate. Fully connected means all-reduce traffic needs no switch hop, but it also means the bandwidth between any one pair is a seventh of the total. Collectives that spread traffic across all seven links can approach the aggregate; a pipeline-parallel stage that only talks to its neighbour sees one link. Prefer tensor parallelism within the node and keep point-to-point-heavy schemes, such as naive expert routing, aware of that per-pair figure. Check the topology with amd-smi topology or rocm-smi --showtopo before trusting a placement.
Settings that decide your numbers
A handful of host and framework settings decide whether you see the data sheet. AMD's tuning guide recommends disabling automatic NUMA balancing, which otherwise migrates pages under running jobs. In PyTorch, TunableOp benchmarks candidate GEMM implementations for your exact shapes and caches the winners. FP8 on MI300X uses the FNUZ formats, so checkpoints quantised to OCP E4M3 must be converted, not reinterpreted.
# Host: disable NUMA auto-balancing (should then print 0)
sudo sh -c 'echo 0 > /proc/sys/kernel/numa_balancing'
cat /proc/sys/kernel/numa_balancing
# PyTorch: tune GEMMs for the shapes your model actually uses
export PYTORCH_TUNABLEOP_ENABLED=1
# Python: the FP8 dtypes MI300X matrix cores consume
import torch
w_fp8 = w_bf16.to(torch.float8_e4m3fnuz) # not torch.float8_e4m3fn
Failure modes
- Assuming a 32 MB L2. Tile orders tuned for a unified cache lose reuse across XCDs. Symptom: a GEMM at a fraction of peak with high fabric traffic in the profiler. Remap program ids or use the library GEMM.
- Hard-coded warp size 32. Shuffles and ballots silently cover half a wavefront. Symptom: wrong reductions, not crashes. Grep ported code for 32 and 0xffffffff.
- Under-filled grids. 304 CUs need several workgroups each. Small-batch decode kernels that launch a few dozen workgroups leave most dies idle; use split-K or stream-K, or a CPX partition for small models.
- FP8 format mismatch. Loading OCP FP8 weights as FNUZ, or the reverse, gives a model that runs and produces garbage, because the exponent bias differs.
- Untuned GEMMs. Default heuristics can miss badly for unusual shapes. Enable TunableOp in a warm-up run and ship the resulting tuning file with the deployment.
Trade-offs
| Choice | Gain | Cost |
|---|---|---|
| SPX, one big device | Simple programming model, all 192 GB in one address space | Round-robin dispatch defeats naive L2 reuse |
| CPX + NPS4 | Memory locality, 5-10% more stream bandwidth, isolation | Eight smaller devices; large models must shard |
| Hand-written Triton with remap | Control over tiling and fusion | You own tuning across shapes and ROCm releases |
| Library GEMMs plus TunableOp | Vendor-tuned kernels with little effort | Less freedom to fuse surrounding ops |
| FP8 FNUZ inference | Twice the matrix throughput, half the weight bytes | Format conversion and accuracy validation |
What to do next
- Run
rocminfo,amd-smi static --partitionandamd-smi topologyon your node and record the gfx target, partition mode and link layout in the runbook. - Disable NUMA auto-balancing and enable TunableOp in a warm-up job; save the tuning results with the deployment.
- For each custom Triton GEMM or attention kernel, add the program-id remap and compare throughput and fabric traffic before and after.
- Grep ported kernels for warp size 32 assumptions and 32-bit lane masks.
- Work out the decode bound for your model: weight bytes divided by 5.3 TB/s. If measured latency is far above it, the gap is software, not hardware.
- Decide SPX or CPX per workload: one large model per GPU in SPX, many small replicas in CPX with NPS4.