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

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:

  • ​row_scales (RowScalesT): The per-row scales.
  • ​first_row (UInt32): Output row of the tile's first row.
  • ​row_end (UInt32): Rows at or past this bound are not loaded and get 0.
  • ​lane (UInt32): Lane index within the warp.
  • ​scale (Float32): Per-expert scale folded into every factor.

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.

Was this page helpful?