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
ScalesLoader
struct ScalesLoader[tma_origin: ImmOrigin, dtype: DType, tile_layout: TensorLayout, desc_layout: TensorLayout = tile_layout, /, *, cta_group: Int]
TMA scales loader parameterized on new Layout types.
Uses TmaOpType to derive the TMATensorTile type from new Layout. Uses async_copy (no multicast). Coordinate order is (row_coord, k_coord) matching scales tensor layout.
Parameters
- tma_origin (
ImmOrigin): Origin of the TMA descriptor pointer. - dtype (
DType): Element data type. - tile_layout (
TensorLayout): Layout of the scales tile loaded into shared memory. - desc_layout (
TensorLayout): Layout of the TMA descriptor (defaults totile_layout). - cta_group (
Int): CTA group size (1 or 2 for SM100 2-SM MMA).
Fields
- tma_op (
ScalesLoader[tma_origin, dtype, tile_layout, desc_layout, cta_group=cta_group].TmaOpPtr):
Implemented traits
AnyType,
Copyable,
Deinitable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
TmaOp
comptime TmaOp = TMATensorTile[dtype, tile_layout.rank, _to_index_list[tile_layout](), _to_index_list[tile_layout.rank, desc_layout]()]
TmaOpPtr
comptime TmaOpPtr = Pointer[TMATensorTile[dtype, tile_layout.rank, _to_index_list[tile_layout](), _to_index_list[tile_layout.rank, desc_layout]()], tma_origin]
Methods
__init__
def __init__[tma_op_type: AnyType](tma_op: Pointer[tma_op_type, tma_origin]) -> Self
Accepts any TMA pointer. Rebinds to the loader's derived type.
Parameters:
- tma_op_type (
AnyType): Compile-time type of the passed TMA descriptor pointer.
Args:
- tma_op (
Pointer[tma_op_type, tma_origin]): Pointer to the TMA descriptor.
load
def load[LayoutType: TensorLayout](self, dest: TileTensor[dtype, LayoutType, MutAnyOrigin, address_space=AddressSpace.SHARED], ref[AddressSpace._value] barrier: SharedMemBarrier, row_coord: Int, k_coord: Int)
Load scales using TMA async copy.
Parameters:
- LayoutType (
TensorLayout): Layout type of the destination TileTensor.
Args:
- dest (
TileTensor[dtype, LayoutType, MutAnyOrigin, address_space=AddressSpace.SHARED]): Destination SMEM TileTensor tile for scales. - barrier (
SharedMemBarrier): Memory barrier for TMA completion signaling. - row_coord (
Int): Row coordinate in global memory (elements). - k_coord (
Int): K dimension coordinate in global memory (elements).