For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/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,
ImplicitlyCopyable,
ImplicitlyDeletable,
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, _safe=True]
Methodsβ
__init__β
def __init__[tma_op_type: AnyType](tma_op: Pointer[tma_op_type, tma_origin, _safe=True]) -> 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, _safe=True]): 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).
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!