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

PreshuffledScaleLoader

struct PreshuffledScaleLoader[MN_padded: Int, K_SCALES: Int]

Per-lane packed-Int32 scale loader from preshuffled GMEM.

Each i32 cell holds 4 E8M0 bytes packed in (k_pack, mn_pack) order; the MFMA's opsel byte index selects the right byte per sub-MMA. OOB lanes (past MN_padded * K_SCALES) read as zero.

Parameters​

  • ​MN_padded (Int): MN dimension rounded up to 32 (the scale-block stride).
  • ​K_SCALES (Int): K // 32 (one E8M0 byte per 32 FP4 elements).

Fields​

  • ​bc (AMDBufferResource):

Implemented traits​

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

Methods​

__init__​

def __init__(scale_gmem_tile: TileTensor[.uint8, Engine=scale_gmem_tile.Engine, address_space=scale_gmem_tile.address_space, linear_idx_type=scale_gmem_tile.linear_idx_type]) -> Self

Builds the V# from a preshuffled per-expert scale byte buffer.

Args:

load_packed​

def load_packed(self, mn: Int, k_scale: Int) -> Int32

Loads the packed Int32 scale word containing logical (mn, k_scale).

Pass (mn, k_scale) at (mn_pack=0, k_pack=0) (the cell base) and all 4 bytes of the cell come back in the returned i32. The MFMA's opsel then selects the byte for each sub-MMA.

Per-lane usage: mn = warp_mn_off + lane % 16 # mn_lane within block k_scale = k_pair_idx * 8 + (lane // 16) # k_lane within block

Args:

  • ​mn (Int): Logical MN index into the [MN_padded, K_SCALES] scale grid.
  • ​k_scale (Int): Logical K scale index into the [MN_padded, K_SCALES] scale grid.

Returns:

Int32

load_group​

def load_group[GROUP: Int](self, mn_base: Int, k_pair_base: Int) -> SIMD[.uint8, (GROUP * Int(4))]

Loads GROUP consecutive packed-scale atoms with one VMEM op.

Consecutive k_pair atoms are contiguous 256-byte blocks, so one GROUP * 4-byte load per lane tiles GROUP of them across a wave64.

The window is lane-transposed: lane l holds the words load_packed would give lanes GROUP*l .. GROUP*l + GROUP - 1, and the caller must undo that.

Parameters:

  • ​GROUP (Int): Atoms per window; GROUP * 4 must be a legal load width.

Args:

  • ​mn_base (Int): Logical MN index of the window's atom row; must be 16-aligned so the lane term is the whole in-atom offset.
  • ​k_pair_base (Int): First k_pair of the window.

Returns:

SIMD[.uint8, (GROUP * Int(4))]

Was this page helpful?