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
TileRowScales
struct TileRowScales[RowScalesT: RowScales, tile_rows: Int]
One output tile's scaled row factors, spread across a warp's lanes.
Lane l loads rows l, l + 32, ... of the tile once, with coalesced
reads. Each epilogue stage then gathers the factors its fragments need
with warp shuffles, so only the first stage waits on global memory.
Parameters
- RowScalesT (
RowScales): The per-row scales type. - tile_rows (
Int): Number of output rows in the tile.
Fields
- values (
SIMD[.float32, ceildiv(tile_rows, _resolve_warp_size())]): - scale (
Float32):
Implemented traits
AnyType,
Copyable,
Deinitable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
rows_per_lane
comptime rows_per_lane = ceildiv(tile_rows, _resolve_warp_size())
Methods
__init__
def __init__(row_scales: RowScalesT, first_row: UInt32, row_end: UInt32, lane: UInt32, scale: Float32) -> Self
Loads this lane's share of the tile's row factors.
Args:
pairs
def pairs[repeats: Int, stage_row: Int](self, lane: UInt32) -> SIMD[.float32, (Int(2) * repeats)]
Returns the factors of one thread's 16x256b accumulator fragments when the fragment column is the output row (transpose_c).
The thread holds fragment columns (lane % 4) * 2 + 8 * r + j for
r < repeats and j < 2. Fragment elements 4 * r + j and
4 * r + 2 + j both sit in that column.
Parameters:
- repeats (
Int): Number of 8-column repeats per fragment. - stage_row (
Int): Tile row of the stage's fragment column 0.
Args:
- lane (
UInt32): Lane index within the warp.
Returns:
SIMD[.float32, (Int(2) * repeats)]: Element 2 * r + j is the factor of column
(lane % 4) * 2 + 8 * r + j, or the expert scale everywhere
when RowScalesT is disabled.
fragment
def fragment[repeats: Int, stage_row: Int](self, lane: UInt32) -> SIMD[.float32, (Int(4) * repeats)]
Returns one factor per element of a thread's 16x256b accumulator fragment when the fragment column is the output row (transpose_c).
Upper and lower fragments differ only in TMEM row, which is the weight dim here, so one result serves both.
Parameters:
- repeats (
Int): Number of 8-column repeats per fragment. - stage_row (
Int): Tile row of the stage's fragment column 0.
Args:
- lane (
UInt32): Lane index within the warp.
Returns:
SIMD[.float32, (Int(4) * repeats)]: The factors, laid out like the fragment.