Layout IR

Layout IR#

Layouts describe how logical tensor coordinates map to storage and execution axes. See Tensor Layout for the programming model and worked examples.

The primary classes are re-exported from tvm.tirx, but are defined in this module:

class tvm.tirx.layout.Layout
verify_well_formed() → bool

Verify if the layout is well-formed.

Returns:

True if the layout is well-formed, False otherwise

Return type:

bool

size(axis_name: str | None = None)

Get the size of the layout.

Parameters:

axis_name (Optional[str]) – The name of the axis to get the size of. If not provided, the default input size will be returned.

span(axis_name: str | None = None)

Get the span of the layout.

Parameters:

axis_name (Optional[str]) – The name of the axis to get the span of. If not provided, the default span will be returned.

apply(*coord: list[Expr], shape: list[Expr] | None = None) → dict[str, Expr]

Apply the layout on the input coordinate and get the mapped output.

Input cases: - coord is a single element -> will be treated as a 1D coordinate - coord is a list of elements -> will be treated as a multi-dimensional coordinate - shape is provided -> turn the coord with shape into a 1D coordinate - shape is not provided -> use the default shape

Returns:

The mapped output (axis name -> value on the axis)

Return type:

Dict[str, Expr]

apply_to_shape(coord: list[Expr], input_shape: list[Expr]) → list[Expr]

Compute the per-shard value that each shard would take if coord were interpreted against input_shape.

Tries self.group(input_shape) first. On success, each group owns exactly one input_shape entry, so coord[d] can be split within that group’s shard extents (bounds stay local to one input dim — simpler analyzer simplification, no cross-dim complications).

Falls back to FlattenCoord(coord, input_shape) + SplitCoord on self’s raw shard shape when the group call fails (e.g. when input_shape does not align with the layout’s factor boundaries).

Returns a list of length len(self.shard); each entry is the value that shard would iterate.

canonicalize() → Layout

Canonicalize the layout by simplifying and fusing iterators where possible.

Returns:

The canonicalized layout

Return type:

Layout

tile(outer: TileLayout, outer_shape: list[Expr], inner_shape: list[Expr]) → TileLayout | ComposeLayout

Tile the current layout with an outer layout.

Parameters:
  • outer (TileLayout) – The outer layout to tile with

  • outer_shape (List[Expr]) – The shape of the outer layout

  • inner_shape (List[Expr]) – The shape of the inner layout

Returns:

The resulting tiled layout

Return type:

Union[TileLayout, ComposeLayout]

direct_sum(left: TileLayout, left_shape: list[Expr], right_shape: list[Expr]) → TileLayout | ComposeLayout

Direct-sum on the tiling domain (unscaled composition): A + B.

This layout is treated as the right addend B grouped by right_shape. The left layout is treated as A grouped by left_shape. The resulting layout is evaluated over the interleaved domain S_A ⊗ S_B, without span scaling (unlike tiling).

is_tile_inner(tile_layout: TileLayout | ComposeLayout, tiled_shape: list[Expr], inner_shape: list[Expr]) → TileLayout | None

Check if a layout is the inner layout of a tiled layout.

Parameters:
  • tile_layout (Union[TileLayout, ComposeLayout]) – The tiled layout to check

  • tiled_shape (List[Expr]) – The shape of the tiled layout

  • inner_shape (List[Expr]) – The shape of the inner layout

Returns:

The outer layout if it is the inner layout of the tiled layout, None otherwise

Return type:

Optional[TileLayout]

is_tile_outer(tile_layout: TileLayout | ComposeLayout, tiled_shape: list[Expr], outer_shape: list[Expr]) → Layout | None

Check if a layout is the outer layout of a tiled layout.

Parameters:
  • tile_layout (Union[TileLayout, ComposeLayout]) – The tiled layout to check

  • tiled_shape (List[Expr]) – The shape of the tiled layout

  • outer_shape (List[Expr]) – The shape of the outer layout

Returns:

The inner layout if it is the outer layout of the tiled layout, None otherwise

Return type:

Optional[Layout]

is_direct_sum_right(sum_layout: TileLayout | ComposeLayout, interleaved_shape: list[Expr], right_shape: list[Expr]) → TileLayout | None

Check if this layout is the right addend B in a direct-sum A + B.

Returns the left addend A if recognized, otherwise None.

is_direct_sum_left(sum_layout: TileLayout | ComposeLayout, interleaved_shape: list[Expr], left_shape: list[Expr]) → Layout | None

Check if this layout is the left addend A in a direct-sum A + B.

Returns the right addend B if recognized, otherwise None.

slice(shape: list[Expr], region: list[tuple[Expr, Expr]]) → Layout | None

Slice the layout with a given shape and region.

Parameters:
  • shape (List[Expr]) – The shape of the layout

  • region (List[Tuple[Expr, Expr], tvm.ir.Range]) – The region to slice, each element is (begin, end)

Returns:

The sliced layout, or None if slicing is not possible

Return type:

Optional[Layout]

tile_to(to_shape: list[Expr], current_shape: list[Expr]) → Layout

Tile the current layout to the given shape.

Parameters:
  • to_shape (List[Expr]) – The shape to tile to

  • current_shape (List[Expr]) – The current shape of the layout

is_swizzle() → bool

Check if the layout is a bare swizzle (ComposeLayout over a trivial tile).

is_trivial() → bool

Check if the layout is trivial.

unpack(num: int) → Layout

Unpack the layout, where a single element in the layout is unpacked into num contiguous elements.

Parameters:

num (int) – The number of elements to unpack into

Returns:

The unpacked layout

Return type:

Layout

broadcast(num: int, position: int = -1, axis: 'Axis' | str = 'm') → Layout

Insert a stride-0 broadcast dim of extent num at position.

position follows Python list-insert semantics (negative indices count from the end; -1 appends after the last shard dim). The new dim has stride 0 — accessing along it doesn’t move the byte offset, so the same physical element is “seen” num times.

Useful for layouts where a consumer reads the same SMEM datum multiple times (e.g. sf_reuse over MMA-K steps).

pack(num: int) → Layout

Pack the layout, where num contiguous elements in the layout are packed into a single element.

Parameters:

num (int) – The number of elements to pack into

Returns:

The packed layout

Return type:

Layout

class tvm.tirx.layout.Axis(*args: Any, **kwargs: Any)

Layout axis wrapper.

property name: str

Return the canonical axis name.

classmethod get(name: str) → Axis

Get or create the axis singleton named name.

Unknown names are registered without thread or memory attributes.

is_thread() → bool

Check if the axis is a thread axis.

is_memory() → bool

Check if the axis is a memory axis.

property backend: str | None

Backend owning this axis, or None for a generic axis.

get_scope() → ExecScope | None

Get the scope of the axis.

get_subscope() → ExecScope | None

Get the subscope of the axis.

class tvm.tirx.layout.Iter(extent: Expr, stride: Expr, axis: Axis | str)

A memory layout that tiles data across devices.

class tvm.tirx.layout.TileLayout(spec: _LayoutSpec)

A memory layout that tiles data across devices.

static from_iters(shard: Sequence[Iter] = (), replica: Sequence[Iter] = (), offset: dict[Axis | str, Expr] | None = None) → TileLayout

Construct a TileLayout from pre-built Iter objects.

is_trivial() → bool

Check if the layout is trivial.

group(shape: list[Expr]) → tuple[Layout, list[int]]

Group the current layout by the given shape.

Parameters:

shape (List[Expr]) – The shape to group by

Returns:

The grouped layout and the separators

Return type:

Tuple[Layout, List[int]]

group_many(shapes: Sequence[Sequence[Expr]]) → tuple[TileLayout, list[list[int]]]

Group the layout by the minimal common refinement of several shapes.

Repeated cumulative product boundaries are retained, so an extent-one dimension in any input shape becomes a real unit iterator in the refined layout. This operation only splits existing shard iterators; it does not canonicalize or reorder the layout.

Parameters:

shapes (Sequence[Sequence[Expr]]) – Logical shapes with provably equal total products.

Returns:

The commonly refined layout and one separator list per input shape.

Return type:

Tuple[TileLayout, List[List[int]]]

get_scope() → tuple[ExecScope, ExecScope] | None

Get the scope pair of the layout.

permute_dims(perm: list[int]) → TileLayout

Permute the dimensions of the layout.

permute_by_groups(seps: list[int], perm: list[int]) → TileLayout

Permute groups of shard iters defined by seps.

seps follows the convention of group()’s second return value: seps[0] == 0 and group i covers shard indices [seps[i], seps[i + 1]). The number of groups is len(seps) - 1.

Parameters:
  • seps (list[int]) – Group boundary positions in the shard list.

  • perm (list[int]) – Permutation of range(len(seps) - 1) selecting the new group order.

class tvm.tirx.layout.ComposeLayout(per_element: int, swizzle_len: int, atom_len: int, tile_layout: TileLayout, swizzle_inner: bool = True)

A memory layout that swizzles a tile layout.

per_element / swizzle_len / atom_len / swizzle_inner carry the swizzle (formerly the standalone SwizzleLayout); tile_layout is the tiled memory map the swizzle is applied to. A bare swizzle is a ComposeLayout over a trivial identity tile.

S[...] and R[...] build shard and replica layout specifications, respectively. Named axes such as laneid, warpid, tid_in_wg, TLane, and TCol belong to CUDA. Use T.cuda.Axis.laneid or import axes from tvm.backend.cuda.axis. Trainium’s P, F, and Bank belong to T.trn.Axis. The generic memory axis remains T.Axis.m.

CUDA memory/instruction layout factories are defined in tvm.backend.cuda.layout and exposed as T.cuda.layout. Trainium PF annotations are explicit factories, for example T.trn.layout.trainium_layout("PF", shape); annotation strings are not accepted by generic tensor constructors.

Omitting layout or passing layout=None leaves every storage scope without a layout. layout="default" explicitly requests a compact TileLayout. Ordinary views preserve shape, strides and element offsets without creating a layout. A tile instruction can infer a temporary layout for a compact global tensor, using the full tensor shape before slicing its operand region. Noncompact global tensors and other storage scopes require an explicit layout when the instruction uses one.

Definition of layout.