Trainium Authoring and Support APIs

Trainium Authoring and Support APIs#

NKI authoring namespace#

The Trainium backend installs Tx.nki. Its current operations are load, store, tensor_copy, matmul, activation, activation_reduce, reciprocal, tensorreduce, tensortensor, tensorscalar, tensorscalar_reduce, scalar_tensor_tensor, scalar_tensor_scalar, memset, identity, and affine_select.

class tvm.backend.trn.script.NKINamespace(op_wrapper: Callable[[Callable[[...], Any]], Callable[[...], Any]] | None = None)

The NKI instructions submodule.

Layout helpers#

Trainium-specific TIRx layout helpers.

tvm.backend.trn.layout.is_trainium_layout(layout: Layout | None) bool

Return whether a layout uses only Trainium memory axes.

tvm.backend.trn.layout.to_psum_layout(layout: TileLayout) TileLayout

Convert a Trainium sbuf layout to its psum physical-bank layout.

tvm.backend.trn.layout.trainium_layout(annotation: str, shape: tuple[Expr], is_psum: bool = False) TileLayout

Create a Trainium tile layout from a PF annotation string and logical shape.

Compilation pipeline#

Trainium TIRX pipeline entrypoints.

tvm.backend.trn.pipeline.finalize_device_passes_trn()

The finalization passes for the Trainium backend.

tvm.backend.trn.pipeline.trn_pipeline()

The Trainium pipeline used in tvm.tirx.build.

Transforms#

Trainium-specific TIRX transformations.

tvm.backend.trn.transform.LowerTIRx()

Lower TIRx tile primitive calls for the Trainium backend.

tvm.backend.trn.transform.LowerTrainiumLayout()

Lower Trainium layouts to backend physical buffer shapes and indices.

class tvm.backend.trn.transform.TrnNaiveAllocator(*args, **kwargs)
class tvm.backend.trn.transform.TrnPrivateBufferAlloc(*args, **kwargs)

Generate private buffer allocations for each TilePrimitiveCall