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

PropertyMI300XWhat it means for software
Compute dies8 XCDs, 304 active CUs (38 per XCD)Grids need well over 304 workgroups to fill the device
L2 cache4 MB per XCD, private to that XCDReuse only happens between workgroups on the same XCD
Infinity Cache256 MB, sharedCatches L2 misses before HBM; large working sets still benefit
HBM3192 GB, 5.3 TB/s peakA 70B model in BF16 fits on one GPU
LDS64 KB per CUBounds tile size and pipeline depth in shared-memory kernels
Peak dense FP16/BF161307.4 TFLOPSMatrix-core figure, not vector ALU
Peak dense FP82614.9 TFLOPSUses the FNUZ FP8 variants, not OCP FP8
Board power750 WSustained clocks depend on cooling and power caps
Peer links128 GB/s per GPU pair, 896 GB/s aggregate in an 8-GPU platformTensor 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

MI300X as kernels see it: private L2 per XCD, shared Infinity Cache, then HBM3XCD 038 CUsL2 4 MBXCD 138 CUsL2 4 MBXCD 238 CUsL2 4 MBXCD 338 CUsL2 4 MBXCD 438 CUsL2 4 MBXCD 538 CUsL2 4 MBXCD 638 CUsL2 4 MBXCD 738 CUsL2 4 MB4 I/O dies: Infinity Fabric network + 256 MB Infinity Cache (shared by all XCDs)192 GB HBM3, 5.3 TB/s peakEach CU: 64 KB LDS, wave64 execution, MFMA matrix units. 8 x 38 = 304 active CUs.An L2 hit stays on one XCD; a miss crosses the fabric to Infinity Cache or HBM.
Data path from a compute unit to HBM. Every level except LDS and L2 is shared device-wide.

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.

SPX mode: workgroup ids are dealt to XCDs round-robinDefault: pid p runs on XCD p mod 80X01X12X23X34X45X56X67X7Tiles 0-7 share one A panel; it ispulled into all eight L2 slicesRemapped: logical tile = (p mod 8) * per + p div 8pids 0, 8, 16, ... 56XCD 0 computes logical tiles 0-7Neighbouring tiles share operandsinside one 4 MB L2Remapping changes which tile a workgroup computes, not where it runs.
Remapping the program id gives each XCD a contiguous block of logical tiles.

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.

PlacementL2 fills per K stepData moved into L2
Round-robin, 1 tile per XCD8 XCDs x (1 A + 1 B)16 slices, 512 KB
Remapped, 1 x 8 tiles on one XCD1 A + 8 B9 slices, 288 KB
Remapped, 2 x 4 block on one XCD2 A + 4 B6 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 NPS4

CPX 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

ChoiceGainCost
SPX, one big deviceSimple programming model, all 192 GB in one address spaceRound-robin dispatch defeats naive L2 reuse
CPX + NPS4Memory locality, 5-10% more stream bandwidth, isolationEight smaller devices; large models must shard
Hand-written Triton with remapControl over tiling and fusionYou own tuning across shapes and ROCm releases
Library GEMMs plus TunableOpVendor-tuned kernels with little effortLess freedom to fuse surrounding ops
FP8 FNUZ inferenceTwice the matrix throughput, half the weight bytesFormat conversion and accuracy validation

What to do next

  1. Run rocminfo, amd-smi static --partition and amd-smi topology on your node and record the gfx target, partition mode and link layout in the runbook.
  2. Disable NUMA auto-balancing and enable TunableOp in a warm-up job; save the tuning results with the deployment.
  3. For each custom Triton GEMM or attention kernel, add the program-id remap and compare throughput and fabric traffic before and after.
  4. Grep ported kernels for warp size 32 assumptions and 32-bit lane masks.
  5. 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.
  6. Decide SPX or CPX per workload: one large model per GPU in SPX, many small replicas in CPX with NPS4.
Key takeaway: MI300X pairs 192 GB of HBM3 with eight compute dies that each own a private 4 MB L2, backed by a shared 256 MB Infinity Cache. In SPX mode workgroups are dealt to the dies round-robin, so kernels that rely on neighbouring tiles sharing L2 should remap program ids to give each XCD a contiguous block. Use wave64-safe code, 16 x 16 MFMA and modest pipeline depth, disable NUMA balancing, tune GEMMs with TunableOp, convert FP8 to FNUZ, and use the capacity to batch decode.