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 function
load_AB_cuda_core
def load_AB_cuda_core[a_type: DType, b_type: DType, a_dim0: Int, a_dim1: Int, a_num_tiles: Int, a_swizzle_bytes: Int, b_dim0: Int, b_dim1: Int, b_num_tiles: Int, b_swizzle_bytes: Int, num_pipeline_stages: Int, //, *, K_actual: Int, cta_group: Int = Int(1), a_swizzle: TensorMapSwizzle = TensorMapSwizzle.SWIZZLE_32B, b_swizzle: TensorMapSwizzle = TensorMapSwizzle.SWIZZLE_32B, a_gmem_layout: Layout = Layout.row_major(Int(1), Int(1)), b_gmem_layout: Layout = Layout.row_major(Int(1), Int(1)), a_plane_splits: IndexList[Int(2)] = Index[Int, Int](Int(0), Int(0))](a_gmem: LayoutTensor[a_type, a_gmem_layout, ImmutAnyOrigin], b_gmem: LayoutTensor[b_type, b_gmem_layout, ImmutAnyOrigin], expert_ids: Pointer[Int32, _safe=False], a_smem_tiles: SMemTileArray2D[a_type, a_dim0, a_dim1, a_num_tiles, a_swizzle_bytes], b_smem_tiles: SMemTileArray2D[b_type, b_dim0, b_dim1, b_num_tiles, b_swizzle_bytes], mma_mbar: Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False], tma_mbar: Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False], producer_phase: PipelineState[num_pipeline_stages], peer_cta_coord: Tuple[Int, Int, Int], work_tile_coord: Tuple[Int, Int], iter_idx: UInt32, scheduler: TileScheduler[static_MN=scheduler.static_MN, tile_shape=scheduler.tile_shape, cluster=scheduler.cluster, cta_group=scheduler.cta_group, swizzle=scheduler.swizzle, swapAB=scheduler.swapAB], qkv_plane_stride: Int = Int(0))
CUDA core fallback for load_AB when K*sizeof < 16 bytes.
Copies [BM, BK] and [BN, BK] tiles from gmem LayoutTensors into swizzled smem, zero-filling columns where k >= K_actual.
Parameters:
- βa_type (
DType): Element type of the A operand tiles (inferred). - βb_type (
DType): Element type of the B operand tiles (inferred). - βa_dim0 (
Int): Number of rows in one A shared-memory tile (inferred). - βa_dim1 (
Int): Number of columns in one A shared-memory tile (inferred). - βa_num_tiles (
Int): Number of pipeline stages in the A shared-memory tile array (inferred). - βa_swizzle_bytes (
Int): Swizzle granularity in bytes for A shared-memory tiles (inferred). - βb_dim0 (
Int): Number of rows in one B shared-memory tile (inferred). - βb_dim1 (
Int): Number of columns in one B shared-memory tile (inferred). - βb_num_tiles (
Int): Number of pipeline stages in the B shared-memory tile array (inferred). - βb_swizzle_bytes (
Int): Swizzle granularity in bytes for B shared-memory tiles (inferred). - βnum_pipeline_stages (
Int): Number of double-buffered pipeline stages for A and B shared-memory tiles (inferred). - βK_actual (
Int): Actual K dimension in elements; columns wherek >= K_actualare zero-filled. - βcta_group (
Int): Number of CTAs cooperating per MMA along the M dimension (defaults to 1). - βa_swizzle (
TensorMapSwizzle): TMA swizzle mode applied to A shared-memory tiles (defaults toSWIZZLE_32B). - βb_swizzle (
TensorMapSwizzle): TMA swizzle mode applied to B shared-memory tiles (defaults toSWIZZLE_32B). - βa_gmem_layout (
Layout): Layout of the A global-memory tensor (defaults toLayout.row_major(1, 1)). - βb_gmem_layout (
Layout): Layout of the B global-memory tensor (defaults toLayout.row_major(1, 1)). - βa_plane_splits (
IndexList[Int(2)]): Per-plane split sizes for fused LoRA QKV A-plane row offsetting;(0, 0)disables it (defaults to(0, 0)).
Args:
- βa_gmem (
LayoutTensor[a_type, a_gmem_layout, ImmutAnyOrigin]): A operand tensor in global memory, source of the[BM, BK]A tiles. - βb_gmem (
LayoutTensor[b_type, b_gmem_layout, ImmutAnyOrigin]): B operand tensor in global memory, source of the[BN, BK]B tiles. - βexpert_ids (
Pointer[Int32, _safe=False]): Pointer to the per-group expert id array; used to offset the A global-memory slice by the expert's local M base. - βa_smem_tiles (
SMemTileArray2D[a_type, a_dim0, a_dim1, a_num_tiles, a_swizzle_bytes]): Shared-memory tile array of staged A tiles, one per pipeline stage. - βb_smem_tiles (
SMemTileArray2D[b_type, b_dim0, b_dim1, b_num_tiles, b_swizzle_bytes]): Shared-memory tile array of staged B tiles, one per pipeline stage. - βmma_mbar (
Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False]): Pointer to the MMA mbarrier array, one per pipeline stage, waited on to confirm the consumer has freed the prior stage's shared memory. - βtma_mbar (
Pointer[SharedMemBarrier, address_space=AddressSpace.SHARED, _safe=False]): Pointer to the TMA mbarrier array, one per pipeline stage, signaled when the copy for a stage is complete. - βproducer_phase (
PipelineState[num_pipeline_stages]): Producer pipeline state tracking the stage index and phase bit. - βpeer_cta_coord (
Tuple[Int, Int, Int]):(peer_id, mma_coord_m, mma_coord_n)tuple giving the peer CTA's coordinates within the cluster. - βwork_tile_coord (
Tuple[Int, Int]):(m, n)element coordinates of the current work tile returned by the scheduler. - βiter_idx (
UInt32): Current K-dimension iteration index within the work tile. - βscheduler (
TileScheduler[static_MN=scheduler.static_MN, tile_shape=scheduler.tile_shape, cluster=scheduler.cluster, cta_group=scheduler.cta_group, swizzle=scheduler.swizzle, swapAB=scheduler.swapAB]): Tile scheduler providing the current group index and static M/N bounds. - βqkv_plane_stride (
Int): Row stride between fused QKV planes used to compute the A-plane row offset (defaults to 0).
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!