IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Mojo struct

TileTensor

struct TileTensor[mut: Bool, //, dtype: DType, LayoutType: TensorLayout, origin: Origin[mut=mut], *, Engine: TensorEngine = DefaultEngine, address_space: AddressSpace = .GENERIC, linear_idx_type: DType = _get_index_type[LayoutType](address_space)]

A tensor type with trait-based layouts supporting nested and hierarchical indexing.

TileTensor provides a flexible abstraction for multi-dimensional data with layouts expressed via the TensorLayout trait. Unlike LayoutTensor which uses a concrete Layout type, TileTensor accepts any type implementing TensorLayout, enabling more flexible compile-time layout composition.

When to use TileTensor vs LayoutTensor:

  • Use TileTensor when you need trait-based layout composition, nested layouts, or when working with the newer Coord-based layout system.
  • Use LayoutTensor when you need established operations like tiled_iterator() or simd_tile(), or compatibility with existing code using IntTuple-based layouts.
  • Both types can interoperate via to_layout_tensor().

Example:

from layout.tile_layout import row_major
from layout import TileTensor
from layout import Idx

# Create a 4x4 tensor with row-major layout
var storage = Array[Float32, 16](uninitialized=True)
var tensor = TileTensor(storage, row_major[4, 4]()).fill(0.0)

# Access elements using flat indices
tensor[0, 0] = 1.0
tensor[1, 2] = 2.0

# Extract a 2x2 tile at position (1, 0)
var tile = tensor.tile[2, 2](1, 0)

# Vectorize for SIMD operations (shape becomes 4x1, element size 1x4)
var vec = tensor.vectorize[1, 4]()

Parameters​

  • ​mut (Bool): The inferred mutability of the underlying pointer.
  • ​dtype (DType): The data type of tensor elements (e.g., DType.float32).
  • ​LayoutType (TensorLayout): A type implementing TensorLayout that defines the tensor's shape and stride structure. Common types include Layout (with Coord-based shapes/strides) and RowMajorLayout.
  • ​origin (Origin[mut=mut]): The origin of the underlying pointer for lifetime tracking.
  • ​Engine (TensorEngine): A type implementing TensorEngine that supplies the storage handle and the load/store/offset operations acting on it. Defaults to DefaultEngine[element_width=1], a plain Pointer handle over non-vectorized elements.
  • ​address_space (AddressSpace): Memory address space (GENERIC, SHARED, CONSTANT, etc.). Defaults to GENERIC.
  • ​linear_idx_type (DType): Integer type for memory indexing. Defaults to int32 for shared/constant memory, int64 otherwise.

Fields​

  • ​layout (LayoutType): The layout instance defining shape and stride mappings.

Implemented traits​

AnyType, Copyable, Deinitable, DevicePassable, ImplicitlyCopyable, Movable, RegisterPassable, TrivialRegisterPassable, Writable

comptime members​

AddressSpaceCastType​

comptime AddressSpaceCastType[address_space: AddressSpace] = TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Type alias for address-space-cast result tensors.

Parameters​

  • ​address_space (AddressSpace): The address_space for the result tensor.

all_dims_known​

comptime all_dims_known = LayoutType.all_dims_known

True if both shape and stride are fully known at compile time.

Required for operations like vectorize() and distribute().

CoalescedType​

comptime CoalescedType = TileTensor[dtype, Layout[TypeList[ComptimeInt[Coord[LayoutType._shape_types].static_product]](), TypeList[ComptimeInt[Int(1)]]()], origin, Engine=Engine, address_space=address_space]

Type alias for coalesced (flattened to rank-1) tensor types.

The coalesced tensor has:

  • shape: product of all original dimensions
  • stride: 1 (contiguous)
  • element shape: product of all original element dimensions
  • element stride: 1 (contiguous)

device_type​

comptime device_type = TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Device-side type for GPU kernel parameter passing.

DeviceGenericType​

comptime DeviceGenericType[origin: Origin[mut=mut]] = TileTensor[dtype, LayoutType, origin, Engine=DevicePointerEngine, linear_idx_type=linear_idx_type]

Type alias for this tensor backed by DevicePointerEngine.

Used by the DeviceBuffer and DevicePointer constructors, which carry the buffer's DevicePointer (its owning reference plus offset and size) to the kernel boundary instead of a bare pointer.

Parameters​

  • ​origin (Origin[mut=mut]): The pointer origin for the returned device-pointer-backed tensor.

DynamicSplitType​

comptime DynamicSplitType[axis: Int = Int(0)] = TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] Scalar[linear_idx_type] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(_mlir_origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type]

Type alias for runtime-sized split element tensors.

The result has an immutable origin.

Parameters​

  • ​axis (Int): The axis along which the tensor is split.

DynamicType​

comptime DynamicType[dyn_dtype: DType] = TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] Scalar[dyn_dtype])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] Scalar[dyn_dtype])]()], origin, Engine=Engine, address_space=address_space]

Type alias for dynamic tensor types.

Parameters​

  • ​dyn_dtype (DType): The data type for Scalar values in the dynamic tensor.

element_size​

comptime element_size = Engine.element_size

Number of scalar elements per logical element, derived from Engine.

ElementType​

comptime ElementType = SIMD[dtype, Engine.element_size]

The SIMD type used for element access.

For scalar tensors, this is SIMD[dtype, 1] (equivalent to Scalar[dtype]). For vectorized tensors, this reflects the vector width.

flat_rank​

comptime flat_rank = TypeList[#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))].length

The flattened rank - total number of dimensions after flattening nested Coords.

For non-nested layouts, flat_rank == rank. For nested layouts (e.g., from blocked_product), flat_rank > rank.

GenericType​

comptime GenericType = TileTensor[dtype, LayoutType, origin, linear_idx_type=linear_idx_type]

Type alias for this tensor with GENERIC address space.

Used by constructors that create tensors from Span or HostBuffer, which produce GENERIC address space tensors.

Immut​

comptime Immut = TileTensor[dtype, LayoutType, origin_of(_mlir_origin), Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Type alias for an immutably-casted tensor.

is_compatible_with​

comptime is_compatible_with[C: TypeList[values]] = ParameterList.all[comptime[Elt: Bool] Elt]() if (Int(len(LayoutType.__shape_types)) == Int(len(values))) else (Int(len(LayoutType.__shape_types)) == Int(len(values)))

True if coordinate types C are structurally compatible with this tensor's layout shape.

A scalar coordinate element is always compatible. A tuple coordinate element requires the corresponding layout shape element to also be a tuple of the same length, checked recursively up to 4 levels of nesting.

Parameters​

is_row_major​

comptime is_row_major = ParameterList.all[comptime[Elt: Bool] Elt]()

True if the tensor has row-major (contiguous) strides.

OffsetViewType​

comptime OffsetViewType[offsets: TypeList[values], LayoutType: TensorLayout = LayoutType] = TileTensor[dtype, LayoutType, origin, Engine=Engine.OffsetResultType[values, offsets], address_space=address_space]

The TileTensor type produced by offsetting into this tensor's storage.

Names the return type of offset-producing operations (slicing, tiling, distribution). It preserves dtype, origin, and address_space, optionally changes the layout, and carries the engine's OffsetResultType[offsets] so an offset that yields a different storage handle is reflected in the view's type.

Parameters​

  • ​offsets (TypeList[values]): The coordinate types of the offset applied to the storage.
  • ​LayoutType (TensorLayout): The layout type of the resulting view. Defaults to this tensor's LayoutType.

OriginCastType​

comptime OriginCastType[mut: Bool, //, origin: Origin[mut=mut]] = TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Type alias for origin-cast result tensors.

Parameters​

  • ​mut (Bool): Whether the result tensor is mutable.
  • ​origin (Origin[mut=mut]): The origin for the result tensor.

rank​

comptime rank = LayoutType.rank

The number of dimensions in the tensor's layout.

ReshapedType​

comptime ReshapedType[*new_shape_types: CoordLike] = TileTensor[dtype, Layout[new_shape_types, TypeList[#kgen.param_list.reduce(#kgen.param_list.tabulate(len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), [idx: __mlir_type.index] values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values[(add (mul idx, -1), len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), -1)]), base=, reducer=[PrevV: KGENParamList[CoordLike], VA: KGENParamList[CoordLike], idx: __mlir_type.index] #kgen.param_list.concat(ComptimeInt[Int(1)] if identical(idx, 0) else Scalar[#kgen.param_list.tabulate(len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), [idx: __mlir_type.index] values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values[(add (mul idx, -1), len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), -1)])[(add idx, -1)].DTYPE if (xor #kgen.param_list.tabulate(len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), [idx: __mlir_type.index] values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values[(add (mul idx, -1), len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), -1)])[(add idx, -1)].is_static_value, True) else PrevV[0].DTYPE] if (xor #kgen.param_list.tabulate(len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), [idx: __mlir_type.index] values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values[(add (mul idx, -1), len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), -1)])[(add idx, -1)].is_static_value, True) if (xor #kgen.param_list.tabulate(len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), [idx: __mlir_type.index] values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values[(add (mul idx, -1), len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), -1)])[(add idx, -1)].is_static_value, True) else (xor PrevV[0].is_static_value, True) else ComptimeInt[Int((mul PrevV[0].static_value, #kgen.param_list.tabulate(len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), [idx: __mlir_type.index] values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values[(add (mul idx, -1), len(values[0]._ParamListType if values[0].is_tuple if identical(len(values), 1) else identical(len(values), 1) else values), -1)])[(add idx, -1)].static_value))], PrevV))]()], origin, Engine=Engine, address_space=address_space]

Type alias for reshaped tensor types.

Parameters​

  • ​*new_shape_types (CoordLike): The shape types for the reshaped tensor.

shape_known​

comptime shape_known = LayoutType.shape_known

True if all shape dimensions are compile-time constants.

SIMDVectorizedType​

comptime SIMDVectorizedType = TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(Int((add LayoutType.__shape_types[idx].static_value, ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value, -1)) // ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=DefaultEngine[element_width=Coord[ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()]].static_product], address_space=address_space, linear_idx_type=linear_idx_type]

Result type for SIMD-width vectorization.

SplitElementType​

comptime SplitElementType[count: Int, axis: Int = Int(0)] = TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // count)] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(_mlir_origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type]

Type alias for equal-sized split element tensors.

The result has an immutable origin.

Parameters​

  • ​count (Int): The number of equal-sized partitions.
  • ​axis (Int): The axis along which the tensor is split.

static_shape​

comptime static_shape[i: Int] = LayoutType.static_shape[i]

Get the compile-time shape value for dimension i, or -1 if dynamic.

Parameters​

  • ​i (Int): The dimension index.

static_stride​

comptime static_stride[i: Int] = LayoutType.static_stride[i]

Get the compile-time stride value for dimension i, or -1 if dynamic.

Parameters​

  • ​i (Int): The dimension index.

StaticSplitType​

comptime StaticSplitType[count: Int, axis: Int = Int(0)] = StaticTuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // count)] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(_mlir_origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type], count]

Type alias for static split result tuples.

Each tuple element is an immutable view.

Parameters​

  • ​count (Int): The number of equal-sized partitions.
  • ​axis (Int): The axis along which the tensor is split.

stride_known​

comptime stride_known = LayoutType.stride_known

True if all stride dimensions are compile-time constants.

TileResultType​

comptime TileResultType[tile_shape_types: TypeList[values], *, linear_idx_type: DType = .int] = TileTensor[dtype, Layout[tile_shape_types, TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]

Result type of .tile[]. Per outer mode, the result stride is parent's innermost sub-stride (CuTe local_tile): identity for scalar parent strides, last-sub-element for tuple parent strides.

Trade-off: the _NestedTileResultStrideTypes[Self.LayoutType] wrap is identity-equivalent for flat parents but nominally a different TypeList than parent's literal _stride_types. Cascaded .tile[].tile[].tile[] chains pay one param_list.tabulate(...) wrap per level (~+20% ASAN compile on linalg matmul kernels with deep cascades). Until Mojo gets dependent return types for parametric aliases (so the flat path could keep parent's literal name) this is the cost of a single unified .tile[] API.

Parameters​

  • ​tile_shape_types (TypeList[values]): The result tile's shape TypeList (typically built from the variadic tile_sizes of the calling .tile[] method).
  • ​linear_idx_type (DType): Integer type keying the offset-derived storage of the result (see OffsetViewType). Defaults to DType.int.

VectorizedType​

comptime VectorizedType[*vector_shape: Int] = TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(Int((add LayoutType.__shape_types[idx].static_value, #kgen.param_list.tabulate(len(values), [idx: __mlir_type.index] ComptimeInt[values[idx]])[idx].static_value, -1)) // #kgen.param_list.tabulate(len(values), [idx: __mlir_type.index] ComptimeInt[values[idx]])[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and #kgen.param_list.tabulate(len(values), [idx: __mlir_type.index] ComptimeInt[values[idx]])[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, #kgen.param_list.tabulate(len(values), [idx: __mlir_type.index] ComptimeInt[values[idx]])[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and #kgen.param_list.tabulate(len(values), [idx: __mlir_type.index] ComptimeInt[values[idx]])[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=DefaultEngine[element_width=Coord[*#kgen.param_list.tabulate(len(values), [idx: __mlir_type.index] ComptimeInt[values[idx]])].static_product], address_space=address_space, linear_idx_type=linear_idx_type]

Type alias for vectorized tensor types.

Parameters​

  • ​*vector_shape (Int): The shape of each vector unit along each axis.

ViewType​

comptime ViewType[new_layout: TensorLayout] = TileTensor[dtype, new_layout, origin, Engine=Engine, address_space=address_space]

A TileTensor type with the same data properties but a different layout.

Preserves dtype, origin, address_space, and other properties while replacing LayoutType. Use this to name the return type of reshape() and other layout-changing operations in helper functions.

Parameters​

  • ​new_layout (TensorLayout): The new TensorLayout type for the view.

Methods​

__init__​

def __init__(var storage: Engine.StorageType[mut, origin_of(origin), dtype, origin, address_space], var layout: LayoutType, /) -> Self

Create a TileTensor from a storage handle and layout.

Args:

  • ​storage (Engine.StorageType[mut, origin_of(origin), dtype, origin, address_space]): The storage handle referencing the tensor data.
  • ​layout (LayoutType): The layout defining the tensor's shape and strides.

def __init__(*, var ptr: Pointer[Scalar[dtype], origin, address_space=address_space], var layout: LayoutType) -> Self

Create a TileTensor from a Pointer and layout.

Args:

def __init__(var span: Span[Scalar[dtype], origin], var layout: LayoutType) -> Self.GenericType

Create a TileTensor from a Span and layout.

Args:

  • ​span (Span[Scalar[dtype], origin]): The memory span containing the tensor data.
  • ​layout (LayoutType): The layout defining the tensor's shape and strides.

Returns:

Self.GenericType

def __init__(ref[origin] device_buffer: DeviceBuffer[dtype], var layout: LayoutType) -> Self.GenericType

Create a LayoutTensor from a DeviceBuffer. The layout must have statically known dimensions.

Note that the device buffer memory is on the accelerator device (GPU global memory). Code running on the CPU can use the DeviceContext to allocate a DeviceBuffer and use that to construct a LayoutTensor that can be accessed on the GPU. You cannot directly access data in the DeviceBuffer or LayoutTensor from the CPU.

The following example shows a typical pattern for using DeviceBuffer to construct a LayoutTensor that you can use on the GPU.

from max.gpu.host import DeviceContext, DeviceBuffer
from layout.tile_layout import row_major
from layout import TileTensor
from layout import Idx

comptime dtype = DType.float32

var ctx = DeviceContext()
# Allocate buffers
var dev_buf = ctx.enqueue_create_buffer[dtype](16)
var host_buf = ctx.enqueue_create_host_buffer[dtype](16)
# Ensure buffers have been created
ctx.synchronize()

# Initialize host buffer and copy to device buffer
for i in range(16):
    host_buf[i] = Scalar[dtype](i)
ctx.enqueue_copy(dev_buf, host_buf)

# Create TileTensor to use on device
var tensor = TileTensor(
     dev_buf,
     row_major(Idx[4], Idx[4]),
)
...

Args:

  • ​device_buffer (DeviceBuffer[dtype]): Contains the underlying data to point to.
  • ​layout (LayoutType): The layout of the tensor.

Returns:

Self.GenericType

def __init__(var device_pointer: DevicePointer[dtype, origin], var layout: LayoutType) -> TileTensor[dtype, LayoutType, origin, Engine=DevicePointerEngine, linear_idx_type=linear_idx_type]

Create a DevicePointerEngine-backed TileTensor from a DevicePointer.

Like the DeviceBuffer constructor, this produces a DevicePointerEngine-backed tile that carries the full DevicePointer (its non-owning reference to the owning DeviceBuffer plus an element offset and size) to the kernel boundary, where DevicePointer._to_device_type encodes it to a bare device pointer. Use this overload when you already hold a DevicePointer (for example an offset one); construct it with TileTensor(buffer.device_ptr(), layout).

The tile borrows the DevicePointer's origin; the backing DeviceBuffer must outlive the tile.

Args:

  • ​device_pointer (DevicePointer[dtype, origin]): The device pointer referencing the tensor data.
  • ​layout (LayoutType): The layout of the tensor.

Returns:

TileTensor[dtype, LayoutType, origin, Engine=DevicePointerEngine, linear_idx_type=linear_idx_type]

def __init__(ref[origin] host_buffer: HostBuffer[dtype], var layout: LayoutType) -> Self.GenericType

Create a LayoutTensor from a HostBuffer. The layout must have statically known dimensions.

The resulting tensor's data can only be accessed on the CPU.

from max.gpu.host import DeviceContext, HostBuffer
from layout.tile_layout import row_major
from layout import TileTensor
from layout import Idx

comptime dtype = DType.float32

var ctx = DeviceContext()
var host_buf = ctx.enqueue_create_host_buffer[dtype](8)

var tensor = TileTensor(
    host_buf,
    row_major(Idx[4], Idx[4]),
)

Args:

  • ​host_buffer (HostBuffer[dtype]): Contains the underlying data to point to.
  • ​layout (LayoutType): The layout of the tensor.

Returns:

Self.GenericType

@implicit def __init__(other: TileTensor[Engine=other.Engine, address_space=other.address_space, linear_idx_type=other.linear_idx_type]) -> other.Immut

Implicitly cast a mutable TileTensor to immutable.

Args:

Returns:

other.Immut

@implicit def __init__(other: TileTensor[Engine=other.Engine, address_space=other.address_space, linear_idx_type=other.linear_idx_type]) -> TileTensor[other.dtype, other.LayoutType, SomeUnsafeAnyOrigin, Engine=other.Engine, address_space=other.address_space, linear_idx_type=other.linear_idx_type]

Implicitly cast a TileTensor to have an AnyOrigin.

Args:

Returns:

TileTensor[other.dtype, other.LayoutType, SomeUnsafeAnyOrigin, Engine=other.Engine, address_space=other.address_space, linear_idx_type=other.linear_idx_type]

__getitem__​

def __getitem__(self, i0: T) -> Self.ElementType

Retrieve the element at the given index or coordinate.

Args:

  • ​i0 (T): The index along axis 0, or a Coord holding every index.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__(self, i0: T, i1: T) -> Self.ElementType

Retrieve the element at the given indices.

Args:

  • ​i0 (T): The index along axis 0.
  • ​i1 (T): The index along axis 1.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__(self, i0: T, i1: T, i2: T) -> Self.ElementType

Retrieve the element at the given indices.

Args:

  • ​i0 (T): The index along axis 0.
  • ​i1 (T): The index along axis 1.
  • ​i2 (T): The index along axis 2.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__(self, i0: T, i1: T, i2: T, i3: T) -> Self.ElementType

Retrieve the element at the given indices.

Args:

  • ​i0 (T): The index along axis 0.
  • ​i1 (T): The index along axis 1.
  • ​i2 (T): The index along axis 2.
  • ​i3 (T): The index along axis 3.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__(self, i0: T, i1: T, i2: T, i3: T, i4: T) -> Self.ElementType

Retrieve the element at the given indices.

Args:

  • ​i0 (T): The index along axis 0.
  • ​i1 (T): The index along axis 1.
  • ​i2 (T): The index along axis 2.
  • ​i3 (T): The index along axis 3.
  • ​i4 (T): The index along axis 4.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__(self, i0: T, i1: T, i2: T, i3: T, i4: T, i5: T) -> Self.ElementType

Retrieve the element at the given indices.

Args:

  • ​i0 (T): The index along axis 0.
  • ​i1 (T): The index along axis 1.
  • ​i2 (T): The index along axis 2.
  • ​i3 (T): The index along axis 3.
  • ​i4 (T): The index along axis 4.
  • ​i5 (T): The index along axis 5.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__[*CoordLikes: CoordLike](self, *coords: *CoordLikes.values) -> Self.ElementType

Retrieve a single element from the tensor at the specified coordinates.

Accepts either a single Coord argument or multiple scalar CoordLike arguments packed into a Coord. Passing a slice produces a view instead.

The fixed-arity overloads above serve ranks up to six; this pack is the fallback beyond that. Overload resolution ranks a fixed arity over any pack, which is what keeps an element load unambiguous against the slicing pack below inside a comptime if branch that is not taken, where where clauses are ignored.

Parameters:

  • ​*CoordLikes (CoordLike): The types of each index argument (CoordLike).

Args:

  • ​*coords (*CoordLikes.values): The coordinates specifying the element's position.

Returns:

Self.ElementType: The element at the specified position.

def __getitem__[*CoordLikes: CoordLike](self, coords: Tuple[*CoordLikes.values]) -> Self.ElementType

Retrieve a single element from the tensor at the specified coordinates.

Accepts either a single Coord argument or multiple scalar CoordLike arguments packed into a Coord.

Parameters:

  • ​*CoordLikes (CoordLike): The types of each index argument (CoordLike).

Args:

Returns:

Self.ElementType: The element at the specified position.

def __getitem__[*arg_types: AnyType](self, *args: *arg_types.values) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])[idx] if (xor conforms_to(#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])[idx], CoordLike), True) else ))), [idx: __mlir_type.index] Scalar[linear_idx_type])](), TypeList[#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx] if (xor conforms_to(arg_types.values[idx], CoordLike), True) else ))]()], origin, Engine=Engine.OffsetResultType[ComptimeInt[_subscript_static_offset[TypeList[#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] _IndexOrSlice[arg_types.values[idx].static_value if conforms_to(arg_types.values[idx], CoordLike) else Int(-1)])](), LayoutType]()], Scalar[linear_idx_type], TypeList[ComptimeInt[_subscript_static_offset[TypeList[#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] _IndexOrSlice[arg_types.values[idx].static_value if conforms_to(arg_types.values[idx], CoordLike) else Int(-1)])](), LayoutType]()], Scalar[linear_idx_type]]()], address_space=address_space] where (Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])[idx] if (xor conforms_to(#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])[idx], CoordLike), True) else )))) > Int(0))

Fix or narrow each dimension, returning a view.

Applies when at least one argument is a slice; indexing every dimension loads a single element instead. Each argument is either an index (n / Idx[n]), which fixes that dimension to the given element and drops it from the result, or a slice (a:b), which keeps the dimension and narrows it to that subrange. Every dimension takes exactly one argument.

Inside a comptime if branch that is not taken, overload resolution ignores where clauses, so this pack alone cannot be told apart from a variadic element overload there. The fixed-arity element overloads above outrank any pack by arity instead, and a conditional return type is no way out either: the compiler's IR verifier does not fold it when an argument's type is an element of a generic Coord.

Compile-time and runtime arguments go through the same path; what is known at compile time is folded there. An Idx[n] index on an axis with a compile-time stride contributes to the view's offset as a ComptimeInt component, and only the remainder is computed at runtime. Slice bounds are runtime values, so a sliced dimension's extent is runtime too -- : is 0:dim, not a marker.

Note: Only works with flat (non-nested) layouts where every shape and stride element is a scalar CoordLike (e.g., layouts produced by row_major, col_major, or manual Layout construction). Does not support nested/hierarchical layouts (e.g., from blocked_product) where shape or stride elements are Coord tuples.

Example:

from layout import TileTensor
from layout.tile_layout import row_major

# 4D tensor: (batch=2, N=8, heads=4, head_dim=16)
var storage = Array[Float32, 2 * 8 * 4 * 16](fill=0)
var t = TileTensor(storage, row_major[2, 8, 4, 16]())

var batch = 1

# Fix batch and heads, keep N and head_dim -> 2D (8, 16). The
# heads index is compile-time, so its share of the offset is too.
var selected = t[batch, :, Idx[2], :]

# Narrow N to a runtime subrange, keep heads and head_dim whole
# -> 3D (n, 4, 16)
var head = t[batch, 0:n, :, :]

Parameters:

  • ​*arg_types (AnyType): The type of each argument: a CoordLike index, or a ContiguousSlice for a dimension to keep.

Args:

  • ​*args (*arg_types.values): One argument per dimension, in dimension order.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])[idx] if (xor conforms_to(#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] arg_types.values[idx])[idx], CoordLike), True) else ))), [idx: __mlir_type.index] Scalar[linear_idx_type])](), TypeList[#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx] if (xor conforms_to(arg_types.values[idx], CoordLike), True) else ))]()], origin, Engine=Engine.OffsetResultType[ComptimeInt[_subscript_static_offset[TypeList[#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] _IndexOrSlice[arg_types.values[idx].static_value if conforms_to(arg_types.values[idx], CoordLike) else Int(-1)])](), LayoutType]()], Scalar[linear_idx_type], TypeList[ComptimeInt[_subscript_static_offset[TypeList[#kgen.param_list.tabulate(len(arg_types.values), [idx: __mlir_type.index] _IndexOrSlice[arg_types.values[idx].static_value if conforms_to(arg_types.values[idx], CoordLike) else Int(-1)])](), LayoutType]()], Scalar[linear_idx_type]]()], address_space=address_space]: A strided view over the same backing storage, of rank equal to the number of slice arguments. Strides are inherited from the surviving axes. The view's offset has a ComptimeInt component for the compile-time indices and a Scalar component for the rest.

__setitem__​

def __setitem__(self, coord: Coord, value: SIMD[dtype, Engine.element_size]) where mut

Set a single element in the tensor at the specified coordinates.

Accepts Coords of flat_rank (flattened).

Args:

def __setitem__[*IndexTypes: Indexer & Copyable](self, *items: *IndexTypes.values, *, value: SIMD[dtype, Engine.element_size]) where ((Int(len(IndexTypes.values)) == Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))))) & mut)

Sets a single element in the tensor at the specified indices.

Uses flat indexing based on flat_rank. For non-nested layouts, flat_rank == rank, so tensor[i, j, k] = value works normally. For nested layouts (e.g., from blocked_product), use all flat_rank indices: tensor[i0, i1, i2, i3] = value for a tensor with flat_rank == 4.

Parameters:

  • ​*IndexTypes (Indexer & Copyable): The types of the indices (must implement Indexer).

Args:

__iadd__​

def __iadd__(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

Adds rhs into this tensor elementwise, in place.

Args:

__isub__​

def __isub__(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

Subtracts rhs from this tensor elementwise, in place.

Args:

__imul__​

def __imul__(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

Multiplies this tensor by rhs elementwise, in place.

Args:

__itruediv__​

def __itruediv__(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

True-divides this tensor by rhs elementwise, in place.

Args:

__ifloordiv__​

def __ifloordiv__(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

Floor-divides this tensor by rhs elementwise, in place.

Args:

get_type_name​

static def get_type_name() -> String

Gets the name of the host type (the one implementing this trait).

Returns:

String: The host type's name.

unsafe_ptr​

def unsafe_ptr(self) -> Pointer[Scalar[dtype], origin, address_space=address_space]

Returns a raw scalar pointer to the base of the tensor's storage.

Delegates to the engine's unsafe_ptr, so the pointer refers to the first scalar element the engine exposes; a vectorized engine (element_size > 1) still yields the scalar base. The pointer borrows the tensor's storage and does not extend its lifetime. Element loads and stores on the tensor go through this pointer.

Returns:

Pointer[Scalar[dtype], origin, address_space=address_space]: A Pointer to Scalar[dtype] at the base of the storage.

load​

def load[width: SIMDLength = Engine.element_size, alignment: Int = Int((get_alignof SIMD[dtype, width], _current_target())) if CompilationTarget.is_gpu() else align_of[dtype](), invariant: Bool = _default_invariant[mut](), non_temporal: Bool = False](self, coord: Coord) -> SIMD[dtype, width]

Load elements from the tensor at the specified coordinates.

Supports both hierarchical indexing (rank indices) and flat indexing (flat_rank indices) for nested layouts.

Parameters:

  • ​width (SIMDLength): Number of elements to load (default: element_size).
  • ​alignment (Int): Memory alignment for the load.
  • ​invariant (Bool): If True, the compiler may assume the memory won't be modified during the kernel, enabling load hoisting and caching.
  • ​non_temporal (Bool): If True, indicates the data will not be reused soon, allowing the hardware to bypass caches (e.g., streaming loads).

Args:

  • ​coord (Coord): The coordinates specifying the element's position.

Returns:

SIMD[dtype, width]: A SIMD vector containing the loaded elements.

store​

def store[width: SIMDLength = Engine.element_size, alignment: Int = Int((get_alignof SIMD[dtype, width], _current_target())) if CompilationTarget.is_gpu() else align_of[dtype](), non_temporal: Bool = False](self, coord: Coord, value: SIMD[dtype, width]) where mut

Store elements to the tensor at the specified coordinates.

Supports both hierarchical indexing (rank indices) and flat indexing (flat_rank indices) for nested layouts.

Parameters:

  • ​width (SIMDLength): Number of elements to store (default: element_size).
  • ​alignment (Int): Memory alignment for the store.
  • ​non_temporal (Bool): If True, indicates the data will not be reused soon, allowing the hardware to bypass caches (e.g., streaming stores).

Args:

  • ​coord (Coord): The coordinates specifying the element's position.
  • ​value (SIMD[dtype, width]): The SIMD vector to store.

load_linear​

def load_linear[width: SIMDLength = Engine.element_size, alignment: Int = Int((get_alignof SIMD[dtype, width], _current_target())), invariant: Bool = _default_invariant[mut]()](self, idx: IndexList[element_type=idx.element_type]) -> SIMD[dtype, width]

Load elements using an IndexList index (for flat layouts).

This enables TileTensor to be used directly with _elementwise_impl_gpu callbacks which pass IndexList coordinates.

Parameters:

  • ​width (SIMDLength): Number of elements to load.
  • ​alignment (Int): Memory alignment for the load.
  • ​invariant (Bool): If True, enables load hoisting.

Args:

Returns:

SIMD[dtype, width]: A SIMD vector containing the loaded elements.

store_linear​

def store_linear[width: SIMDLength = Engine.element_size, alignment: Int = Int((get_alignof SIMD[dtype, width], _current_target()))](self: TileTensor[dtype, Engine=self.Engine, address_space=self.address_space, linear_idx_type=self.linear_idx_type], idx: IndexList[element_type=idx.element_type], value: SIMD[dtype, width])

Store elements using an IndexList index (for flat layouts).

This enables TileTensor to be used directly with _elementwise_impl_gpu callbacks which pass IndexList coordinates.

Parameters:

  • ​width (SIMDLength): Number of elements to store.
  • ​alignment (Int): Memory alignment for the store.

Args:

raw_load​

def raw_load[width: SIMDLength = 1, alignment: Int = align_of[dtype](), invariant: Bool = _default_invariant[mut](), non_temporal: Bool = False](self, offset: T) -> SIMD[dtype, width]

Load width elements starting at ptr[offset], bypassing the layout.

This is a raw read against the underlying storage: the caller is responsible for ensuring offset is a valid index into the backing buffer, independent of the tensor's layout. Useful for kernels that treat the buffer as a contiguous array (copies, fills, reductions over contiguous storage).

Parameters:

  • ​width (SIMDLength): Number of elements to load.
  • ​alignment (Int): Memory alignment for the load.
  • ​invariant (Bool): If True, enables load hoisting.
  • ​non_temporal (Bool): If True, indicates the data will not be reused soon, allowing the hardware to bypass caches (e.g., streaming loads).

Args:

  • ​offset (T): Linear element offset into the underlying storage.

Returns:

SIMD[dtype, width]: A SIMD vector containing the loaded elements.

raw_store​

def raw_store[width: SIMDLength = 1, alignment: Int = align_of[dtype](), non_temporal: Bool = False](self, offset: T, value: SIMD[dtype, width]) where mut

Store width elements at ptr[offset], bypassing the layout.

This is a raw write against the underlying storage: the caller is responsible for ensuring offset is a valid index into the backing buffer, independent of the tensor's layout.

Parameters:

  • ​width (SIMDLength): Number of elements to store.
  • ​alignment (Int): Memory alignment for the store.
  • ​non_temporal (Bool): If True, indicates the data will not be reused soon, allowing the hardware to bypass caches (e.g., streaming stores).

Args:

  • ​offset (T): Linear element offset into the underlying storage.
  • ​value (SIMD[dtype, width]): The SIMD vector to store.

as_span​

def as_span(self) -> Span[Scalar[dtype], origin, address_space=address_space]

Get a Span over the tensor's elements.

Constraints:

The tensor must have row-major (contiguous) strides, so storage order and layout order agree and the span visits every element exactly once.

Returns:

Span[Scalar[dtype], origin, address_space=address_space]: A Span of num_elements() scalars over the tensor's storage.

bitcast​

def bitcast[target_dtype: DType](self) -> TileTensor[target_dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Reinterprets the tensor's element dtype, preserving layout.

Returns a new TileTensor that shares the same underlying storage and layout as self but views elements as target_dtype rather than Self.dtype.

Parameters:

  • ​target_dtype (DType): The new element dtype to view the storage as.

Returns:

TileTensor[target_dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]: A TileTensor[target_dtype, ...] backed by the same pointer and layout as self.

ptr_at_offset​

def ptr_at_offset(self, coords: Coord) -> Pointer[Scalar[dtype], origin, address_space=address_space] where (Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType)))) == Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))))) or (Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType)))) == Int(1))

Get a pointer offset at the given flattened coordinates.

Args:

  • ​coords (Coord): A flattened list of the offset coordinates.

Returns:

Pointer[Scalar[dtype], origin, address_space=address_space]: A pointer offset at the given flattened coordinates.

prefetch​

def prefetch(self, coords: Coord) where (Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(coords.element_types.values), [idx: __mlir_type.index] coords.element_types.values[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType)))) == Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType)))))

Prefetch tensor data at the specified coordinates into cache.

Issues a software prefetch hint to the processor to load the data at coords into the cache hierarchy. This can improve performance by reducing memory latency for subsequent accesses to the same location.

Performance:

  • Prefetching is a performance hint and does not guarantee data will be cached.
  • Most effective when issued sufficiently ahead of the actual data access.
  • Uses high locality prefetch to the data cache, optimized for data that will be accessed multiple times.
  • Can reduce memory access latency by 50-90% when used correctly.

Notes:

  • Excessive prefetching can pollute the cache and degrade performance.
  • Most beneficial for predictable access patterns that would otherwise cause cache misses.
  • No operation is performed on the prefetched data.

Args:

  • ​coords (Coord): The indices.

num_elements​

def num_elements(self) -> Int

Returns the total number of elements in the tensor.

Computes the product of all shape dimensions.

Returns:

Int: The total element count.

copy_from​

def copy_from(self, other: TileTensor[Engine=other.Engine, address_space=other.address_space, linear_idx_type=other.linear_idx_type]) where mut

Copy data from another tensor into this tensor.

Performs an element-by-element copy from other into self, respecting the layouts of both tensors. Each logical element is loaded from other using its layout and stored into self using self's layout, so the copy works correctly even when the tensors have different shapes or strides (as long as they agree on total element count).

When both tensors have fully static, row-major layouts, the copy widens to SIMD load + cast + SIMD store, using the narrower of the two dtypes' native SIMD widths.

The copy loop lives in the engines. This forwards self and other as (storage, layout) pairs to the source engine's copy_to, which by default hands them to Self.Engine.copy_from; an engine whose storage has no pointer, such as TMemEngine, overrides copy_to to run its own load loop.

  • Both tensors must have statically known shapes with matching total element count.
  • Source and destination dtypes may differ; each logical element is cast to the destination dtype.

Args:

copy_from_async​

def copy_from_async[is_masked: Bool = False, swizzle: Optional[Swizzle] = None, fill: Fill = Fill.NONE, eviction_policy: CacheEviction = CacheEviction.EVICT_NORMAL](self, src: TileTensor[Engine=src.Engine, address_space=src.address_space, linear_idx_type=src.linear_idx_type], src_idx_bound: Scalar[src.linear_idx_type] = 0, base_offset: Scalar[linear_idx_type] = 0) where mut

Asynchronously copy data from another tensor to this tensor using GPU hardware.

This method performs an asynchronous copy from the source tensor to this tensor using GPU hardware acceleration. It's specifically designed for copying data from global memory to shared memory in GPU kernels, leveraging hardware-specific asynchronous copy mechanisms for improved performance.

For optimal performance, you need to arrange the copy correctly. Use the distribute() method to create thread-local fragments of the source and destination tensors, assigning each thread one or more elements to copy.

Optionally, use the vectorize() method to get vectorized views of both tensors before calling distribute(). This allows each thread to copy multiple elements of the tensor. For example:

var fragment = tensor.vectorize[1, simd_width]().distribute[
    thread_layout
](thread_id)

The copy operation is asynchronous, so you must call async_copy_wait_all() or async_copy_wait_group() to ensure the copy has completed before using the data.

Unlike LayoutTensor, a TileTensor's logical element is always a contiguous run of element_size scalars (the engine's element_width), so there is no non-vectorizable element layout to fall back on: every copy issues one cp.async per logical element.

Example:

from layout import Idx, TileTensor, row_major
from layout.tile_tensor import stack_allocation
from max.gpu import thread_idx
from max.gpu.memory import async_copy_commit_group, async_copy_wait_all
from max.gpu.sync import barrier

def kernel(src_ptr: MutPointer[Float32, MutAnyOrigin]):
    comptime thread_layout = row_major(Idx[2], Idx[2])

    var src = TileTensor(src_ptr, row_major[4, 4]())
    var smem = stack_allocation[
        dtype = DType.float32, address_space = .SHARED
    ](row_major[4, 4]())

    # Each of the 4 threads copies its own 2x2 fragment.
    var tid = thread_idx.x
    smem.distribute[thread_layout](tid).copy_from_async(
        src.distribute[thread_layout](tid)
    )
    async_copy_commit_group()
    async_copy_wait_all()
    barrier()
    # ... read the shared tile

Performance:

  • Supports vectorized copies for 4, 8, or 16-byte elements for better throughput.
  • Can bypass L1 cache with appropriate eviction policies for specific access patterns.
  • Swizzling can improve memory access patterns and reduce bank conflicts.

Notes:

  • Asynchronous copies allow computation to overlap with memory transfers.
  • A synchronization barrier is required before using the copied data.

Constraints:

  • Destination must be in shared memory.
  • Source must be in the generic or global address space.
  • Source and destination data types must match.
  • Element size must be 4, 8, or 16 bytes.
  • Destination tensor must have a static layout.
  • Fill.NAN requires a floating-point dtype and 16-byte elements.

Parameters:

  • ​is_masked (Bool): Whether to perform a masked copy, where elements outside the src_idx_bound are not copied and are handled according to fill instead.
  • ​swizzle (Optional[Swizzle]): Optional swizzling function to rearrange the destination indices, which can improve memory access patterns.
  • ​fill (Fill): What a masked copy does with the bytes it skips. Fill.NONE leaves them untouched, so the destination keeps whatever it already held. Fill.ZERO zeroes them, byte-granularly, so a partially valid element is part copy and part zero. Fill.NAN writes NaN, but only whole elements at a time: a partially valid element is filled rather than partly copied.
  • ​eviction_policy (CacheEviction): Cache eviction policy for the source data.

Args:

write_to​

def write_to(self, mut w: T)

Format and write the tensor's contents to a writer.

Uses bracket-delimited, comma-separated format. For 2D tensors, the output shows nested row structure. For other ranks, values are printed as a flat bracketed list in column-major coordinate order.

Example:

from layout import TileTensor
from layout.tile_layout import row_major

def main():
    var storage = Array[Float32, 4](uninitialized=True)
    var vec = TileTensor(storage, row_major[4]()).fill(1.0)
    print(vec)   # [1.0, 1.0, 1.0, 1.0]

    var storage2 = Array[Float32, 6](uninitialized=True)
    var mat = TileTensor(storage2, row_major[2, 3]()).fill(1.0)
    print(mat)   # [[1.0, 1.0, 1.0], [1.0, 1.0, 1.0]]

Args:

  • ​w (T): The writer instance to write the formatted output to.

tile​

def tile[*tile_sizes: Int](self, coordinates: Coord) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]

Extract a sub-tile (CuTe local_tile). Works on both flat and nested parent layouts.

On a flat parent, returns the rank-Self.rank sub-tile whose strides are the parent's strides and shape is tile_sizes. On a nested parent of shape ((outer_h, inner_h), (outer_w, inner_w)), slices one outer index per mode and returns a flat rank-2 sub-tile whose strides are each parent mode's innermost sub-strides.

Parameters:

  • ​*tile_sizes (Int): The dimensions of the tile along each axis.

Args:

  • ​coordinates (Coord): The tile coordinates as a Coord.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]: A view into the original tensor representing the sub-tile.

def tile[*tile_sizes: Int, *, stride_layout: TensorLayout](self, coordinates: Coord) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), stride_layout._shape_types], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]

Tile with explicit static strides (flat parents only).

Use when the parent tensor has dynamic (Scalar) strides but the actual stride values are known at compile time. This produces a tile with all_dims_known=True, enabling vectorize/distribute.

This is needed because TensorLayout trait parameters erase concrete stride types -- the compiler cannot prove all_dims_known through a trait-bounded parameter even when the underlying strides are static.

Parameters:

  • ​*tile_sizes (Int): Tile dimensions along each axis.
  • ​stride_layout (TensorLayout): The layout providing static stride types.

Args:

  • ​coordinates (Coord): Tile coordinates in the grid.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), stride_layout._shape_types], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]: A view into the original tensor representing the specified tile.

def tile[tile_shape_types: TypeList[tile_shape_types.values], //](self, tile_shape: Coord[tile_shape_types], coordinates: Coord) -> TileTensor[dtype, Layout[tile_shape_types, TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]

Extract a tile (sub-tensor) with shape specified as a Coord argument.

This overload accepts the tile shape as a Coord value rather than compile-time Int parameters, enabling use cases where tile shapes are constructed programmatically or passed as values.

Example:

from layout.tile_layout import row_major
from layout import TileTensor
from layout.coord import coord

var storage = Array[Float32, 16](uninitialized=True)
var tensor = TileTensor(storage, row_major[4, 4]()).fill(1.0)

# Extract the tile at position (1, 0) with tile size 2x2
var t = tensor.tile(coord[2, 2], coord[1, 0])

Parameters:

Args:

Returns:

TileTensor[dtype, Layout[tile_shape_types, TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]: A view into the original tensor representing the specified tile.

def tile[*tile_sizes: Int](self, *tile_coords: Int) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]

Variadic-Int-coords form of .tile[]. Works on both flat and nested parents: see the Coord-arg sibling above.

Example:

from layout.tile_layout import row_major
from layout import TileTensor

var storage = Array[Float32, 16](uninitialized=True)
var tensor = TileTensor(storage, row_major[4, 4]()).fill(1.0)

# Extract the tile at position (1, 0) with tile size 2x2
var t = tensor.tile[2, 2](1, 0)

Parameters:

  • ​*tile_sizes (Int): The dimensions of each tile along each axis.

Args:

  • ​*tile_coords (Int): The coordinates of the specific tile to extract.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[idx]._ParamListType[(add len(LayoutType.__stride_types[idx]._ParamListType), -1)] if LayoutType.__stride_types[idx].is_tuple else LayoutType.__stride_types[idx])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]: A view into the original tensor representing the sub-tile.

tile_with_offset​

def tile_with_offset[*tile_sizes: Int](self, coordinates: Coord) -> Tuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), LayoutType._stride_types], origin, Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space], IndexList[Int(len(coordinates.element_types.values))], Int]

Like tile(), but also returns corner coordinates and linear offset. Flat-layout parents only.

Parameters:

  • ​*tile_sizes (Int): Tile dimensions along each axis.

Args:

  • ​coordinates (Coord): Tile coordinates in the grid.

Returns:

Tuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), LayoutType._stride_types], origin, Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space], IndexList[Int(len(coordinates.element_types.values))], Int]: Tuple of (tile, corner_coords, offset).

def tile_with_offset[*tile_sizes: Int, *, stride_layout: TensorLayout](self, coordinates: Coord) -> Tuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), stride_layout._shape_types], origin, Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space], IndexList[Int(len(coordinates.element_types.values))], Int]

Like tile(), but with explicit static strides. Flat-layout parents only.

Use when the parent has dynamic strides but the values are known at compile time. See tile[stride_layout=...] for details.

Parameters:

  • ​*tile_sizes (Int): Tile dimensions along each axis.
  • ​stride_layout (TensorLayout): The layout providing static stride types.

Args:

  • ​coordinates (Coord): Tile coordinates in the grid.

Returns:

Tuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(tile_sizes.values), [idx: __mlir_type.index] ComptimeInt[tile_sizes.values[idx]])](), stride_layout._shape_types], origin, Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space], IndexList[Int(len(coordinates.element_types.values))], Int]: Tuple of (tile, corner_coords, offset).

reshape​

def reshape[new_layout: TensorLayout](self, layout_val: new_layout) -> TileTensor[dtype, new_layout, origin, Engine=Engine, address_space=address_space]

Create a view of the tensor with a different layout.

Returns a new TileTensor sharing the same pointer but with a different layout. This is a zero-cost operation -- only the layout type changes, no data is moved.

Parameters:

  • ​new_layout (TensorLayout): The target layout type (inferred from layout_val).

Args:

  • ​layout_val (new_layout): The layout instance to use for the new view.

Returns:

TileTensor[dtype, new_layout, origin, Engine=Engine, address_space=address_space]: A TileTensor with the new layout viewing the same memory.

def reshape[*new_shape: Int](self) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])](), TypeList[#kgen.param_list.reduce(#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)]), base=, reducer=[PrevV: KGENParamList[CoordLike], VA: KGENParamList[CoordLike], idx: __mlir_type.index] #kgen.param_list.concat(ComptimeInt[Int(1)] if identical(idx, 0) else Scalar[#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].DTYPE if (xor #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].is_static_value, True) else PrevV[0].DTYPE] if (xor #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].is_static_value, True) if (xor #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].is_static_value, True) else (xor PrevV[0].is_static_value, True) else ComptimeInt[Int((mul PrevV[0].static_value, #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].static_value))], PrevV))]()], origin, Engine=Engine, address_space=address_space] where TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].all_dims_known and TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].is_row_major and (Coord[LayoutType._shape_types].static_product == Coord[*#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])].static_product)

Reshape the tensor to a new shape with compile-time dimensions.

This method creates a view of the tensor with a different logical shape while preserving the underlying data. The total number of elements must remain the same, and the tensor must have row-major (contiguous) strides.

Example:

from layout.tile_layout import row_major
from layout import TileTensor

var storage = Array[Float32, 12](uninitialized=True)
var tensor = TileTensor(storage, row_major[3, 4]()).fill(1.0)
# tensor has shape (3, 4)

var reshaped = tensor.reshape[2, 6]()
# reshaped has shape (2, 6), same underlying data

var reshaped_1d = tensor.reshape[12]()
# reshaped_1d has shape (12,), equivalent to coalesce

Performance:

  • Creates a view without copying data.
  • Zero-cost abstraction at compile time when used with static shapes.

Constraints:

  • All dimensions must be statically known (all_dims_known).
  • The tensor must have row-major strides (is_row_major).
  • The product of the new shape must equal the product of the original shape.

Parameters:

  • ​*new_shape (Int): The new shape dimensions as compile-time integers.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])](), TypeList[#kgen.param_list.reduce(#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)]), base=, reducer=[PrevV: KGENParamList[CoordLike], VA: KGENParamList[CoordLike], idx: __mlir_type.index] #kgen.param_list.concat(ComptimeInt[Int(1)] if identical(idx, 0) else Scalar[#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].DTYPE if (xor #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].is_static_value, True) else PrevV[0].DTYPE] if (xor #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].is_static_value, True) if (xor #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].is_static_value, True) else (xor PrevV[0].is_static_value, True) else ComptimeInt[Int((mul PrevV[0].static_value, #kgen.param_list.tabulate(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), [idx: __mlir_type.index] #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[(add (mul idx, -1), len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0]._ParamListType if #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])[0].is_tuple if identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else identical(len(#kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), 1) else #kgen.param_list.tabulate(len(new_shape.values), [idx: __mlir_type.index] ComptimeInt[new_shape.values[idx]])), -1)])[(add idx, -1)].static_value))], PrevV))]()], origin, Engine=Engine, address_space=address_space]: A TileTensor with the new shape and row-major strides, sharing the same underlying data as the original tensor.

def reshape[*new_shape_types: CoordLike](self, new_shape: Coord[new_shape_types]) -> TileTensor[dtype, Layout[new_shape_types, TypeList[#kgen.param_list.reduce(#kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)]), base=, reducer=[PrevV: KGENParamList[CoordLike], VA: KGENParamList[CoordLike], idx: __mlir_type.index] #kgen.param_list.concat(ComptimeInt[Int(1)] if identical(idx, 0) else Scalar[#kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].DTYPE if (xor #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].is_static_value, True) else PrevV[0].DTYPE] if (xor #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].is_static_value, True) if (xor #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].is_static_value, True) else (xor PrevV[0].is_static_value, True) else ComptimeInt[Int((mul PrevV[0].static_value, #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].static_value))], PrevV))]()], origin, Engine=Engine, address_space=address_space] where TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].is_row_major

Reshape the tensor to a new shape specified as a Coord.

This method creates a view of the tensor with a different logical shape while preserving the underlying data. The total number of elements must remain the same, and the tensor must have row-major (contiguous) strides.

This overload accepts shapes with runtime dimensions, performing the element count validation at runtime when needed.

Example:

from layout.tile_layout import row_major
from layout import TileTensor
from layout import Idx, Coord

var storage = Array[Float32, 12](uninitialized=True)
var tensor = TileTensor(storage, row_major[3, 4]()).fill(1.0)

# Reshape with runtime-determined dimensions
var rows = 2
var cols = 6
var reshaped = tensor.reshape(Coord(rows, cols))

Performance:

  • Creates a view without copying data.
  • May include runtime validation for dynamic shapes.

Constraints:

  • The tensor must have row-major strides (is_row_major).
  • The product of the new shape must equal the product of the original shape (validated at runtime for dynamic shapes).

Parameters:

  • ​*new_shape_types (CoordLike): The types of the new shape dimensions (inferred).

Args:

Returns:

TileTensor[dtype, Layout[new_shape_types, TypeList[#kgen.param_list.reduce(#kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)]), base=, reducer=[PrevV: KGENParamList[CoordLike], VA: KGENParamList[CoordLike], idx: __mlir_type.index] #kgen.param_list.concat(ComptimeInt[Int(1)] if identical(idx, 0) else Scalar[#kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].DTYPE if (xor #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].is_static_value, True) else PrevV[0].DTYPE] if (xor #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].is_static_value, True) if (xor #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].is_static_value, True) else (xor PrevV[0].is_static_value, True) else ComptimeInt[Int((mul PrevV[0].static_value, #kgen.param_list.tabulate(len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), [idx: __mlir_type.index] new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values[(add (mul idx, -1), len(new_shape_types.values[0]._ParamListType if new_shape_types.values[0].is_tuple if identical(len(new_shape_types.values), 1) else identical(len(new_shape_types.values), 1) else new_shape_types.values), -1)])[(add idx, -1)].static_value))], PrevV))]()], origin, Engine=Engine, address_space=address_space]: A TileTensor with the new shape and row-major strides, sharing the same underlying data as the original tensor.

transpose​

def transpose(self) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[(add (mul idx, -1), len(LayoutType.__shape_types), -1)])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[(add (mul idx, -1), len(LayoutType.__stride_types), -1)])]()], origin, Engine=Engine, address_space=address_space]

Create a transposed view of the tensor.

Returns a new TileTensor sharing the same pointer but with the layout dimensions reversed. For 2D tensors, this swaps rows and columns. This is a zero-cost operation -- no data is moved.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[(add (mul idx, -1), len(LayoutType.__shape_types), -1)])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] LayoutType.__stride_types[(add (mul idx, -1), len(LayoutType.__stride_types), -1)])]()], origin, Engine=Engine, address_space=address_space]: A TileTensor with transposed layout viewing the same memory.

distribute​

def distribute[thread_layout: Layout[thread_layout.shape_types, thread_layout.stride_types], swizzle: Optional[Swizzle] = None](self, thread_id: Int) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // thread_layout.shape_types.values[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, thread_layout.shape_types.values[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]

Distribute tensor workload across multiple threads in a structured pattern.

This method partitions a tensor across multiple threads for parallel processing, assigning each thread a specific portion of the tensor. The distribution pattern is determined by the thread_layout parameter, which defines the logical arrangement of threads.

Parameters:

Args:

  • ​thread_id (Int): The ID of the current thread (0-based).

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // thread_layout.shape_types.values[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, thread_layout.shape_types.values[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=Engine.OffsetResultType[Scalar[linear_idx_type], TypeList[Scalar[linear_idx_type]]()], address_space=address_space]: A view into the original tensor representing the portion assigned to this thread.

distribute_with_offset​

def distribute_with_offset[thread_layout: Layout[thread_layout.shape_types, thread_layout.stride_types], swizzle: Optional[Swizzle] = None](self, thread_id: Int) -> Tuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // thread_layout.shape_types.values[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, thread_layout.shape_types.values[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space], IndexList[Int(len(thread_layout.shape_types.values))], Int]

Like distribute(), but also returns thread coordinates and offset.

Parameters:

Args:

  • ​thread_id (Int): The ID of the current thread (0-based).

Returns:

Tuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // thread_layout.shape_types.values[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, thread_layout.shape_types.values[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and thread_layout.shape_types.values[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space], IndexList[Int(len(thread_layout.shape_types.values))], Int]: Tuple of (distributed_tensor, thread_coords, offset).

fill​

def fill[*, use_runtime_layout: Bool = not TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].all_dims_known.__bool__() if not TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].all_dims_known.__bool__() else (Coord[LayoutType._shape_types].static_product > Int(2048))](self, val: Scalar[dtype]) -> Self where mut

Fill the entire tensor with a single value.

This method sets all elements of the tensor to the specified value. It works with both statically and dynamically shaped tensors.

For statically known layouts, the fill operation is unrolled at compile time. For dynamic layouts, a runtime loop is used. No vectorization is applied, so performance may be suboptimal for large tensors. Consider using hardware-specific fill operations for better performance with large tensors.

This method can be used with tensors of any rank and shape. The fill operation respects the tensor's layout, filling all elements regardless of how they are arranged in memory. For tensors with element_layout, all elements within each logical element are filled with the same value.

Example:

from layout.tile_layout import row_major
from layout import TileTensor

def main() raises:
    var storage = Array[Float32, 3 * 4](uninitialized=True)
    var tensor = TileTensor(storage, row_major[3,4]()).fill(0.0)
    print(tensor)

If not using method chaining, you can either reassign the result to the tensor variable, or assign the result to the discard pattern (_) to avoid warnings about an unused value:

from layout.tile_layout import row_major
from layout import TileTensor

var storage = Array[Float32, 3 * 4](uninitialized=True)
var tensor = TileTensor(storage, row_major[3,4]()).fill(0.0)
tensor = tensor.fill(0.0)
# or
_ = tensor.fill(0.0)

Parameters:

  • ​use_runtime_layout (Bool): Whether to use the runtime layout for filling. This parameter is defaulted to True if the layout is not statically known. If loop bounds are too large, it's better to use the runtime layout to avoid long compilation time.

Args:

  • ​val (Scalar[dtype]): The value to fill the tensor with. Must be of the same data type as the tensor.

Returns:

Self: The tensor itself (self), allowing for method chaining.

dim​

def dim[i: Int](self) -> Scalar[linear_idx_type]

Returns the size of outer-mode dimension i.

For a flat layout this is shape[i]. For a nested layout (where shape[i] is itself a Coord) this is the product of all leaf dims under outer-mode i: the i-th mode's extent under CuTe Layout Algebra. For shape ((a, b), (c, d)): dim[0] = a*b, dim[1] = c*d.

Parameters:

  • ​i (Int): The dimension index (compile-time constant).

Returns:

Scalar[linear_idx_type]: The product of all leaf dims under outer-mode i.

def dim[IndexType: Indexer](self, index: IndexType) -> Scalar[linear_idx_type]

Returns the size of the specified dimension.

Parameters:

  • ​IndexType (Indexer): The type of the index argument.

Args:

  • ​index (IndexType): The dimension index (runtime value).

Returns:

Scalar[linear_idx_type]: The size of the specified dimension as a scalar.

dynamic_stride​

def dynamic_stride[IndexType: Indexer](self, index: IndexType) -> Scalar[linear_idx_type]

Returns the stride of the specified dimension.

Parameters:

  • ​IndexType (Indexer): The type of the index argument.

Args:

  • ​index (IndexType): The dimension index (runtime value).

Returns:

Scalar[linear_idx_type]: The stride of the specified dimension as a scalar.

split​

def split[count: Int, axis: Int = Int(0)](self) -> StaticTuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // count)] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type], count] where not mut

Splits the tensor into equal-sized views along an axis.

Split views are immutable. Call as_imm().split[count]() on a mutable tensor before splitting.

See also: Use split(count, idx) to return a single partition with a runtime-sized split axis. The dynamic overload takes axis before split_alignment as compile-time parameters, while this overload takes count before axis.

Constraints:

The tensor shape must be statically known. The split axis must have static, scalar shape and stride values. The split-axis shape must be evenly divisible by count.

Parameters:

  • ​count (Int): The number of partitions to split into.
  • ​axis (Int): The axis along which to split.

Returns:

StaticTuple[TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(LayoutType.__shape_types[idx].static_value // count)] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type], count]: A StaticTuple containing count non-overlapping TileTensor views into this tensor.

def split[axis: Int = Int(0), split_alignment: Int = Int(1)](self, count: Int, idx: Int) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] Scalar[linear_idx_type] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type] where not mut

Returns one partition of the tensor after splitting along an axis.

The returned partition is immutable. Call as_imm().split(count, idx) on a mutable tensor before splitting.

The base partition size is align_up(ceildiv(axis_dim, count), split_alignment). This can make the first count - 1 partitions larger than ceildiv(axis_dim, count); each returned view is clamped to the remaining elements. If the aligned partition offsets exhaust the axis before all count partitions are assigned, trailing partitions have size 0.

See also: Use split[count]() to split into a StaticTuple of equal-sized views when the partition count is known at compile time.

Parameters:

  • ​axis (Int): The axis along which to split.
  • ​split_alignment (Int): Alignment for the partition size.

Args:

  • ​count (Int): The number of partitions.
  • ​idx (Int): The partition index to return.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] Scalar[linear_idx_type] if identical(idx, axis) else LayoutType.__shape_types[idx])](), LayoutType._stride_types], origin_of(origin), Engine=Engine.OffsetResultType[Int, TypeList[Int]()], address_space=address_space, linear_idx_type=linear_idx_type]: An immutable TileTensor view whose split axis has runtime shape.

slice​

def slice[*slices: _IndexOrSlice](self) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] _indexed_extent[slices, LayoutType, Int(idx)]())), [idx: __mlir_type.index] ComptimeInt[#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] _indexed_extent[slices, LayoutType, Int(idx)]())[idx]])](), TypeList[#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] LayoutType.static_stride[_kept_slice_axis_for_output[slices, Int(idx)]()])), [idx: __mlir_type.index] ComptimeInt[#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] LayoutType.static_stride[_kept_slice_axis_for_output[slices, Int(idx)]()])[idx]])]()], origin, Engine=Engine.OffsetResultType[ComptimeInt[_slice_storage_offset[slices, LayoutType]()], TypeList[ComptimeInt[_slice_storage_offset[slices, LayoutType]()]]()], address_space=address_space] where (Int(len(slices.values)) == Int(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(#kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))), [idx: __mlir_type.index] #kgen.param_list.concat(#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] LayoutType.__shape_types[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))[idx]._ParamListType))))) and TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].all_dims_known

Extract a view of the tensor, fixing or subslicing each dimension.

Each parameter is either an Int, which fixes that dimension to the given index and drops it from the result (rank reduction, like squeeze), or a slice literal selecting a rank-preserving subrange -- e.g. t.slice[0:5, 2](). The output rank equals the number of slice parameters. Unlike tile, whose coordinate indexes a grid of tile-sized blocks, slice bounds are element offsets, so a view need not be aligned to its own extent (e.g. the shorter trailing tile of a tile_iterator walk). The bounds must be compile-time values: they are folded into the view's storage handle as a ComptimeInt offset, so a view of a fully static tensor stays fully static.

Example:

from layout.tile_layout import row_major
from layout import TileTensor

comptime layout_3d = row_major[16, 16, 16]()
var stack = Array[UInt8, layout_3d.static_product](fill=0)
var tensor_3d = TileTensor(stack, layout_3d)

# Rank-preserving: a 2x2x4 view.
var sub = tensor_3d.slice[0:2, 1:3, 0:4]()

# Rank-reducing: plane 3, then a 2x4 view of it.
var plane = tensor_3d.slice[3, 1:3, 0:4]()

Performance:

  • Creates a view without copying data, making it very efficient.
  • Maintains the original tensor's stride information for efficient memory access.
  • Free at runtime: the whole view, offset included, is computed in the type system and emits no index arithmetic.

Notes:

  • The slice is a view into the original tensor, so modifications to the slice will affect the original tensor.
  • The step size must be 1 for all dimensions (no gaps allowed).
  • Slice bounds are checked at compile time against the parent's static shape; out-of-range bounds are a compile error, not a runtime one.

Parameters:

  • ​*slices (_IndexOrSlice): One Int or slice literal per tensor dimension.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] _indexed_extent[slices, LayoutType, Int(idx)]())), [idx: __mlir_type.index] ComptimeInt[#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] _indexed_extent[slices, LayoutType, Int(idx)]())[idx]])](), TypeList[#kgen.param_list.tabulate(len(#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] LayoutType.static_stride[_kept_slice_axis_for_output[slices, Int(idx)]()])), [idx: __mlir_type.index] ComptimeInt[#kgen.param_list.tabulate(_count_slice_dims[slices](), [idx: __mlir_type.index] LayoutType.static_stride[_kept_slice_axis_for_output[slices, Int(idx)]()])[idx]])]()], origin, Engine=Engine.OffsetResultType[ComptimeInt[_slice_storage_offset[slices, LayoutType]()], TypeList[ComptimeInt[_slice_storage_offset[slices, LayoutType]()]]()], address_space=address_space]: A strided sub-view over the same backing storage. Its extents are fresh ComptimeInts, since a sliced extent is a new compile-time value rather than the parent's; its strides are the surviving axes' own.

vectorize​

def vectorize[*vector_shape: Int](self) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(Int((add LayoutType.__shape_types[idx].static_value, #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].static_value, -1)) // #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=DefaultEngine[element_width=Coord[*#kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])].static_product], address_space=address_space, linear_idx_type=linear_idx_type]

Reshape a tensor into a vectorized form for efficient SIMD operations.

This method transforms the tensor's logical layout to enable efficient vectorized processing, treating blocks of elements as vector units. The transformation is particularly useful for SIMD (Single Instruction Multiple Data) operations and hardware acceleration.

The vector shape is tracked in element_size.

Example:

For a 16x16 tensor, vectorize[4, 4] will produce a 4x4 tensor where each element position is the starting point of a 4x4 block from the original tensor. The strides are scaled by the vector shape so that adjacent elements in the vectorized tensor are spaced apart by the vector dimensions.

Performance:

  • Creates a view without copying data, making it very efficient.
  • Enables strided access patterns suitable for SIMD vector loads.
  • Zero-cost abstraction at compile time when used with static shapes.

Parameters:

  • ​*vector_shape (Int): The dimensions of each vector unit along each axis of the tensor. For example, in a 2D tensor, vectorize[4, 4] treats 4x4 blocks as vector units.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(Int((add LayoutType.__shape_types[idx].static_value, #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].static_value, -1)) // #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and #kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])[idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=DefaultEngine[element_width=Coord[*#kgen.param_list.tabulate(len(vector_shape.values), [idx: __mlir_type.index] ComptimeInt[vector_shape.values[idx]])].static_product], address_space=address_space, linear_idx_type=linear_idx_type]: A view of the tensor with a vectorized layout, where each element in the resulting tensor represents the start of a vector block from the original tensor. The element layout is tracked via element_size (the vector shape).

def vectorize(self) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(Int((add LayoutType.__shape_types[idx].static_value, ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value, -1)) // ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=DefaultEngine[element_width=Coord[ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()]].static_product], address_space=address_space, linear_idx_type=linear_idx_type]

Return a SIMD-width vectorized view of this tensor.

This is a convenience method that vectorizes along the last dimension by the SIMD width for the tensor's dtype.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] ComptimeInt[(Int((add LayoutType.__shape_types[idx].static_value, ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value, -1)) // ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value)] if LayoutType.__shape_types[idx].is_static_value and ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].is_static_value else Scalar[LayoutType.__shape_types[idx].DTYPE])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] ComptimeInt[Int((mul LayoutType.__stride_types[idx].static_value, ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].static_value))] if LayoutType.__stride_types[idx].is_static_value and ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()][idx].is_static_value else Scalar[LayoutType.__stride_types[idx].DTYPE])]()], origin, Engine=DefaultEngine[element_width=Coord[ComptimeInt[Int(1)], ComptimeInt[simd_width_of[dtype]()]].static_product], address_space=address_space, linear_idx_type=linear_idx_type]: A Self.VectorizedType[1, simd_width_of[Self.dtype]()] view whose last dimension stride equals the SIMD width for the tensor's dtype.

coalesce​

def coalesce(self) -> Self.CoalescedType where TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].all_dims_known and TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type].is_row_major

Creates a rank-1 tensor by flattening all dimensions.

Coalescing combines all dimensions into a single contiguous dimension. This is useful for operations that need to iterate over all elements sequentially.

Example:

For a 4x4 tensor, coalesce() produces a 16-element rank-1 tensor. For a vectorized tensor with shape (4, 4) and element shape (4, 4), coalescing produces shape (16,) with element shape (16,).

Performance:

  • Creates a view without copying data.
  • Enables simple sequential iteration over all elements.
  • Zero-cost abstraction at compile time.

Constraints:

All dimensions must be statically known (all_dims_known). The tensor must have row-major (contiguous) strides (is_row_major).

Returns:

Self.CoalescedType: A rank-1 tensor with shape equal to the product of all original dimensions and stride 1. Element layout is also coalesced.

make_dynamic​

def make_dynamic[dyn_dtype: DType](self) -> TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] Scalar[dyn_dtype])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] Scalar[dyn_dtype])]()], origin, Engine=Engine, address_space=address_space]

Convert all elements in shape and stride to Scalar[dyn_dtype].

Examples:

from layout import TileTensor
from layout.tile_layout import row_major
var storage = Array[Float32, 12](uninitialized=True)
var tensor = TileTensor(Span(storage), row_major[3, 4]())
var dynamic = tensor.make_dynamic[.int64]()
# dynamic has Int64 for all shape/stride dimensions

Parameters:

  • ​dyn_dtype (DType): The data type for the resulting Scalar values.

Returns:

TileTensor[dtype, Layout[TypeList[#kgen.param_list.tabulate(len(LayoutType.__shape_types), [idx: __mlir_type.index] Scalar[dyn_dtype])](), TypeList[#kgen.param_list.tabulate(len(LayoutType.__stride_types), [idx: __mlir_type.index] Scalar[dyn_dtype])]()], origin, Engine=Engine, address_space=address_space]: A new TileTensor where all elements in shape and stride are converted to Scalar[dyn_dtype].

to_layout_tensor​

def to_layout_tensor(self) -> LayoutTensor[dtype, Layout(coord_to_int_tuple[LayoutType._shape_types](), coord_to_int_tuple[LayoutType._stride_types]()), origin, address_space=address_space]

Return a LayoutTensor with the same shape, stride, and address space of this tensor.

This is a utility to help with porting LayoutTensor methods to this type.

Supports DefaultEngine and DevicePointerEngine-backed tiles. For a DevicePointerEngine-backed tile the raw device pointer is recovered from the handle (via Engine.unsafe_ptr), so the resulting LayoutTensor no longer carries the owning DevicePointer. This is a temporary workaround until LayoutTensor support is removed as part of GPUA-6.

Returns:

LayoutTensor[dtype, Layout(coord_to_int_tuple[LayoutType._shape_types](), coord_to_int_tuple[LayoutType._stride_types]()), origin, address_space=address_space]: A LayoutTensor with the same shape, stride, and address space of this tensor.

as_unsafe_any_origin​

def as_unsafe_any_origin(self) -> TileTensor[dtype, LayoutType, SomeUnsafeAnyOrigin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Casts the origin of the TileTensor to UnsafeAnyOrigin.

Safety:

It is always preferred to maintain a concrete origin values instead of using UnsafeAnyOrigin. Casting to UnsafeAnyOrigin is an inherently unsafe operation that will silently extend unrelated lifetimes and turn off exclusivity checking.

Returns:

TileTensor[dtype, LayoutType, SomeUnsafeAnyOrigin, Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]: A tensor with the origin set to UnsafeAnyOrigin.

as_immut​

def as_immut(self) -> TileTensor[dtype, LayoutType, origin_of(origin), Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Return an immutable version of this tensor.

Deprecated: 'as_immut' is deprecated, use 'as_imm' instead

Returns:

TileTensor[dtype, LayoutType, origin_of(origin), Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]: A TileTensor covering the same elements, but without mutability.

as_imm​

def as_imm(self) -> TileTensor[dtype, LayoutType, origin_of(origin), Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]

Return an immutable version of this tensor.

Returns:

TileTensor[dtype, LayoutType, origin_of(origin), Engine=Engine, address_space=address_space, linear_idx_type=linear_idx_type]: A TileTensor covering the same elements, but without mutability.

address_space_cast​

def address_space_cast[target_address_space: AddressSpace](self) -> TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=target_address_space, linear_idx_type=linear_idx_type]

Return a version of this tensor cast to a new address space.

Parameters:

  • ​target_address_space (AddressSpace): The target address space to cast to.

Returns:

TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=target_address_space, linear_idx_type=linear_idx_type]: A TileTensor covering the same elements in the new address space.

unsafe_address_space_cast​

def unsafe_address_space_cast[target_address_space: AddressSpace](self) -> TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=target_address_space, linear_idx_type=linear_idx_type]

Return a version of this tensor cast to a new address space.

Parameters:

  • ​target_address_space (AddressSpace): The target address space to cast to.

Returns:

TileTensor[dtype, LayoutType, origin, Engine=Engine, address_space=target_address_space, linear_idx_type=linear_idx_type]: A TileTensor covering the same elements in the new address space.

to_device_buffer​

def to_device_buffer(self, ctx: DeviceContext) -> DeviceBuffer[dtype]

Convert the tensor to a DeviceBuffer.

Works for tensors backed by either DefaultEngine or DevicePointerEngine. In both cases the base pointer is recovered through the engine (self.ptr), so the resulting non-owning DeviceBuffer covers exactly this tensor's elements, honoring any offset baked into the storage handle.

Args:

Returns:

DeviceBuffer[dtype]: A DeviceBuffer containing the tensor's data.

min​

def min(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

Takes the elementwise minimum with rhs, in place.

Args:

max​

def max(self, rhs: TileTensor[dtype, Engine=rhs.Engine, address_space=rhs.address_space, linear_idx_type=rhs.linear_idx_type]) where mut and conforms_to(Engine, TensorOps)

Takes the elementwise maximum with rhs, in place.

Args:

abs​

def abs(self) where mut and conforms_to(Engine, TensorOps)

Takes the elementwise absolute value of this tensor, in place.

For unsigned dtypes this is the identity.

recip​

def recip(self) where mut and conforms_to(Engine, TensorOps)

Replaces each element of this tensor with its reciprocal, in place.

Elements equal to zero produce infinity, following IEEE 754 division semantics.

Constraints:

The tensor's dtype must be a floating-point type.

exp​

def exp[scale_dtype: DType = dtype, //, scale: Scalar[scale_dtype] = 1](self) where mut and conforms_to(Engine, TensorOps)

Replaces each element x of this tensor with exp(scale * x), in place.

The scale factor is applied before exponentiation so that scaled exponentials (for example softmax logit scaling) fuse into a single pass over the elements. The default scale of 1 gives a plain exponential.

Constraints:

The tensor's dtype must be a floating-point type.

Parameters:

  • ​scale_dtype (DType): The data type of the scale factor. Defaults to the tensor's dtype; the scale is cast to the tensor's dtype before the multiplication.
  • ​scale (Scalar[scale_dtype]): The compile-time factor each element is multiplied by before exponentiation.

Was this page helpful?