Chapter 4 · Programming models and profiling
Other hardware stacks

Nothing about the tile abstraction is NVIDIA-specific, and the clearest evidence comes from Pallas, JAX’s extension for “kernel programming for both GPUs and TPUs using a Triton-like model.” Its design document makes a pointed observation about Triton itself: Triton “exposes a TPU-like programming model to users, i.e. writing programs for tiles of arrays in L1-cache,” and yet is specialized enough to GPU that it cannot be compiled directly for TPU; Triton’s atomic operations, built for parallel writes, “don’t necessarily make sense on TPU.” So Pallas keeps only the tile-based programming model, abstracts the platform details behind it, and lowers the same kernel per backend: to Mosaic GPU (formerly Triton) on GPUs, and to Mosaic on TPUs. The tile is the portable part; what surrounds it is not.

What Pallas adds around that tile is exactly three things, and the design document is emphatic that it is otherwise “just JAX.” First, reference types: a kernel receives Refs rather than immutable arrays, which “gives users more precise control over memory access and layout,” read and written in place with NumPy-style indexing, and output Refs are handed in alongside the input ones. Reading from a Ref “corresponds to loading an array into the lowest level of the memory hierarchy,” L1 cache on GPU and vector registers on TPU, which is this chapter’s tile wearing a JAX name. Second, a restricted subset of JAX primitives plus Pallas-specific ones, among them a program_id the document admits is borrowed terminology from Triton. Third, pallas_call, the higher-order function that runs the kernel, “analogous to pmap or shard_map, except with references to shared memory,” which loops over a grid and slices each iteration’s inputs through a BlockSpec carrying a block shape and an index map from loop indices to block indices.

The third one is the answer to why a kernel language would be embedded in an array framework rather than built beside it. The stated motivation is that “JAX transformations are key to its success,” so the same vmap and jvp that transform ordinary JAX programs can be applied to a kernel the user wrote by hand. A Triton kernel is a leaf the framework calls into; a Pallas kernel is a JAX program, and the transformations do not stop at its edge.

AMD’s stack arrived at the same shape from two directions. Composable Kernel, in AMD’s own description, “provides a programming model for writing performance-critical kernels for machine learning workloads across multiple architectures,” written in general purpose kernel languages such as HIP C++ and resting on two ideas: “a tile-based programming model” and a complexity-reduction technique it calls . The library names four layers, ordered lowest to highest: Templated Tile Operators, Templated Kernel and Invoker, Instantiated Kernel and Invoker, and Client API. That is the CUTLASS decomposition again under different names, an atom-like tile operator at the bottom and a ready-made handle at the top, and the repository documents two operation sets against it, one for CK and a separate one for CK Tile. HipKittens, a research library of C++ tile primitives for AMD GPUs, states this section’s through-line as an experimental finding: porting from its NVIDIA sibling ThunderKittens, the core tile and bulk compute interfaces carry over, while the decisions around memory access patterns, compute and memory scheduling, and thread block ordering within the chiplet architecture differ. Its tiles are sized to the tensor core units, and its two core scheduling patterns, 8-wave ping pong and 4-wave interleave, are organized around CDNA’s waves rather than the warps of Foundations.

Where those AMD kernels end up is its own data point. AITER, the AI Tensor Engine for ROCm, is “AMD’s high-performance AI operator library, providing optimized GPU kernels for inference and training workloads on ROCm,” positioned as “a unified collection of production-ready operators that framework developers can integrate directly into their stacks.” What makes it relevant here is its backend list: the same operator can be served by a Triton kernel, a Composable Kernel implementation, or hand-tuned assembly, behind one C++ or Python API. The programming models this chapter has covered are not competing endpoints so much as interchangeable suppliers to an operator library, and AITER’s README notes it is the default attention backend for vLLM on ROCm, which is where a kernel chosen by that machinery actually meets production traffic.

AWS Trainium is the strongest test of the claim, because the accelerator underneath is not a GPU. NKI, the Neuron Kernel Interface for writing kernels that run on Trainium devices, still hands the programmer a tile: a kernel allocates tiles in on-chip SBUF memory, DMA-copies inputs into them from HBM, and checks that a tile’s first dimension fits within the on-chip tile size limit before operating on whole tiles at a time. What changes is everything around the tile. NKI kernels use a rather than a grid of parallel blocks, and its nki.isa functions are “designed to expose the underlying hardware capabilities in as direct a way as possible,” each call running one operation on one of the device’s compute engines while the compiler unrolls, inlines, and resolves everything else ahead of time. Across all four stacks the tile survives even where warps, blocks, and threads do not; what each stack builds around the tile tracks what its particular silicon makes cheap or expensive. See further reading for each stack’s documentation.

Inside NKI: two namespaces and a specializer#

The tile survives on Trainium; almost none of the rest of this chapter’s vocabulary does. The shape of a NKI kernel is easiest to see in its two namespaces. nki.language, imported as nl, is the high-level API: tensor creation through ndarray, data types, memory buffers, loop ranges such as affine_range and dynamic_range, and math operations like nl.add, nl.matmul, and nl.softmax, many of which are “convenience wrappers around one or more nki.isa operations.” nki.isa, imported as nisa, is the other end: “each function in this namespace maps directly to a Trainium hardware operation,” and they “are the only calls that produce runtime operations on the device.” The guide states the trade between them in one line: “Use nki.language for readability and portability; use nki.isa when you need precise control over which hardware engine executes an operation.”

What makes that split load-bearing rather than stylistic is , the first of the three stages a @nki.jit kernel goes through, ahead of compilation to Trainium machine code and linking into the Neuron graph compiler’s larger computation graph. During specialization “the compiler acts as an interpreter for the meta-programming parts of your kernel”: every for loop that is not a dynamic_range is unrolled, every function call is inlined, every if on a compile-time condition is resolved to its taken branch, and every Python expression over compile-time values is evaluated away. What comes out is “a specialized, flat sequence of nki.isa.* operations with all compile-time values resolved.”

So the rule a Trainium kernel author reasons with is unusually crisp: “the only constructs that survive specialization and become runtime operations are nki.isa.* calls and dynamic_range loops.” Everything else is meta-programming whose only job is to decide which ISA operations get generated. That is a different contract from Triton’s, and the difference is not a matter of degree. Triton’s compiler is handed decisions the programmer declined to make, the coalescing and swizzling and shared-memory allocation and thread mapping listed at the top of this chapter; NKI’s compiler is handed a program whose runtime content is already a list of hardware operations, and asked to evaluate away everything that is not one. Both languages give you a tile. Only one of them is choosing instructions on your behalf.

The tile still carries a hardware-shaped constraint, and the guide’s introductory kernel checks it before doing anything else: a_input.shape[0] must be no larger than nl.tile_size.pmax, the on-chip memory tile size, and the assertion is there because this particular kernel does not tile its inputs. What follows is the whole of NKI’s memory model in four calls. Tiles are allocated with buffer=nl.sbuf, both inputs are DMA-copied in from HBM with nisa.dma_copy, the addition is a single nisa.tensor_tensor, and the result is DMA-copied back out to nl.shared_hbm:

A NKI kernel body, from the NKI language guide
assert a_input.shape[0] <= nl.tile_size.pmax

a_tile = nl.ndarray(shape=a_input.shape, dtype=a_input.dtype, buffer=nl.sbuf)
nisa.dma_copy(dst=a_tile, src=a_input)

b_tile = nl.ndarray(shape=b_input.shape, dtype=b_input.dtype, buffer=nl.sbuf)
nisa.dma_copy(dst=b_tile, src=b_input)

c_tile = nl.ndarray(shape=a_input.shape, dtype=a_input.dtype, buffer=nl.sbuf)
nisa.tensor_tensor(dst=c_tile, data1=a_tile, data2=b_tile, op=nl.add)

c_output = nl.ndarray(dtype=a_input.dtype, shape=a_input.shape, buffer=nl.shared_hbm)
nisa.dma_copy(dst=c_output, src=c_tile)

Four ISA calls, no grid, no threads, and no launch configuration. The sequential model means those four run in the order they are written, the compiler free to reorder only operations with no data dependency between them, which it may do because the reordering is “functionally transparent to NKI programmers.” Set that beside the Triton vector addition at the top of this chapter, which is also a handful of statements over tiles, and the surface similarity is real. The difference is what each one is a specification of: the Triton kernel specifies a computation and leaves the instruction selection open, while the NKI kernel already is the instruction sequence, written in Python.