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
PreshuffledBLoader
struct PreshuffledBLoader[N: Int, K_BYTES: Int, cache_policy: CacheOperation = CacheOperation.ALWAYS, lane_bytes: Int = Int(16)]
Per-lane B fragment loader from preshuffled GMEM (DRAM -> VGPR direct).
The 5D layout places each lane's 16-byte fragment at a contiguous DRAM
offset, so a single buffer_load_dwordx4 per lane delivers the MFMA's
B operand with no LDS staging. OOB lanes are clamped to zero by the
buffer-resource bounds.
Parameters
- N (
Int): Per-expert N dimension (rows of the logical [N, K_BYTES] tile). - K_BYTES (
Int): Per-expert FP4-packed K dimension (= K // 2). - cache_policy (
CacheOperation): Cache hint for the B load. Defaults toALWAYS(normal cached, flydslb_nt=0); setSTREAMING(NT=1, flydslb_nt=2) to skip caching B fragments that are streamed once and never reused. - lane_bytes (
Int): Bytes one lane feeds the MFMA: 16 for FP4, 24 for FP6, 32 for FP8. Widths above 16, or not a power of two, are split into planes (seeShuffler.b_plane_byte_off) and loaded with one instruction each.
Fields
- bc (
AMDBufferResource):
Implemented traits
AnyType,
Copyable,
Deinitable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
num_planes
comptime num_planes = Shuffler.num_planes[lane_bytes]()
reg_bytes
comptime reg_bytes = Int(16) if (lane_bytes <= Int(16)) else Int(32)
Methods
__init__
def __init__(b_gmem_tile: TileTensor[.uint8, Engine=b_gmem_tile.Engine, address_space=b_gmem_tile.address_space, linear_idx_type=b_gmem_tile.linear_idx_type]) -> Self
Builds the V# from a preshuffled per-expert B byte buffer.
Args:
- b_gmem_tile (
TileTensor[.uint8, Engine=b_gmem_tile.Engine, address_space=b_gmem_tile.address_space, linear_idx_type=b_gmem_tile.linear_idx_type]): Preshuffled per-expert B byte buffer holding the[N, K_BYTES]logical tile, as produced byblock_scaled_preshuffle_layouts.
load_fragment
def load_fragment(self, n: Int, k_byte: Int) -> SIMD[.uint8, Int(16) if (xor identical((lt lane_bytes, 17), False), True) else Int(32)]
Loads one lane's B fragment at logical (n, k_byte).
For one MFMA dispatch a lane calls this with
(n = warp_n_off + n_mma * 16 + lane % 16, k_byte = k_tile * MFMA_K_BYTES + (lane // 16) * lane_bytes).
A single-plane fragment is one buffer_load_dwordx4, byte-identical to
the layout this loader has always used. A multi-plane fragment issues
one naturally-aligned load per plane and assembles them in registers;
the payload stays contiguous, which is what the MFMA requires.
Args:
- n (
Int): Logical N row index into the[N, K_BYTES]tile. - k_byte (
Int): Logical K byte index into the[N, K_BYTES]tile.
Returns:
SIMD[.uint8, Int(16) if (xor identical((lt lane_bytes, 17), False), True) else Int(32)]
lane_plane_off
def lane_plane_off[plane: Int = Int(0)](self, n: Int, lane_k_byte: Int) -> Int32
Returns the K-invariant per-lane part of one plane's address.
Pair with load_at, which supplies the wave-uniform whole-tile part.
Splitting the address this way lets a caller hoist the per-lane term
out of an unrolled K loop instead of rematerialising it per tile.
b_plane_byte_off is additive in k0: every other term depends only
on n and the lane's own K offset, and the k0 * tile_bytes stride is
the same for all planes because tile_bytes is derived from
lane_bytes, not from the plane width. So a multi-plane fragment
(FP6's 16 + 8) splits per plane exactly as a single-plane one does.
Parameters:
- plane (
Int): Which plane of the lane fragment to address.
Args:
- n (
Int): Logical N row index into the[N, K_BYTES]tile. - lane_k_byte (
Int): The lane's own K byte offset within its K tile, i.e.(lane // 16) * lane_bytes.
Returns:
load_at
def load_at[plane: Int = Int(0)](self, lane_off: Int32, k_byte_uniform: Int) -> SIMD[.uint8, Shuffler.plane_bytes[lane_bytes, plane]()]
Loads one plane at lane_off plus a wave-uniform K offset.
k_byte_uniform is a whole number of (n0, k0) tiles, so it rides
soffset while lane_off stays in voffset. The tile stride is a
function of lane_bytes alone, so it is shared by every plane; only
the loaded width narrows on a partial plane (FP6's second is 8 bytes).
Parameters:
- plane (
Int): Which plane of the lane fragment to load.
Args:
- lane_off (
Int32): The per-lane offset fromlane_plane_off. - k_byte_uniform (
Int): Logical K byte base of the tile, wave-uniform and a multiple of the K0 tile width.
Returns: