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
SoftmaxRDNA
struct SoftmaxRDNA[dtype: DType, num_m_mmas: Int, num_n_mmas: Int, num_warps_m: Int, num_warps_n: Int, use_exp2: Bool = False]
Online softmax accumulator for RDNA Wave32 attention kernels.
Tracks per-row running maximum and sum across online softmax iterations
so that attention scores can be normalized incrementally as new key/value
tiles arrive. The full() method runs one complete iteration (max → exp →
sum → correction → output update) end-to-end.
Parameters
- dtype (
DType): Element type of the score and output fragments. - num_m_mmas (
Int): Number of MMA tiles along the query (row) dimension. - num_n_mmas (
Int): Number of MMA tiles along the key (column) dimension. - num_warps_m (
Int): Number of warps assigned to the row dimension. - num_warps_n (
Int): Number of warps assigned to the column dimension. - use_exp2 (
Bool): Selectsexp2instead of naturalexpfor the softmax kernel.
Fields
- rowmax_tensor (
SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].RowMaxTensorType): - rowsum_tensor (
SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].RowSumTensorType): - score_frag_rowmax (
SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].ScoreFragTensorType): - score_frag_rowsum (
SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].ScoreFragTensorType): - correction (
SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].ScoreFragTensorType):
Implemented traits
comptime members
exp_function
comptime exp_function = _exp2_concrete[_, _] if use_exp2 else _exp_concrete[_, _]
frag_is_row_vector
comptime frag_is_row_vector = True
frag_num_rows
comptime frag_num_rows = ComptimeInt[Int(1)].static_value
frag_size
comptime frag_size = SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].FragmentLayoutT.static_product
FragmentLayoutT
comptime FragmentLayoutT = Layout[*(), *()]
num_colwise_lanes
comptime num_colwise_lanes = SIMD(Int(16))
num_colwise_tiles
comptime num_colwise_tiles = num_m_mmas
num_colwise_warps
comptime num_colwise_warps = num_warps_m
num_rowwise_lanes
comptime num_rowwise_lanes = SIMD(Int(2))
num_rowwise_tiles
comptime num_rowwise_tiles = num_n_mmas
num_rowwise_warps
comptime num_rowwise_warps = num_warps_n
row_layout
comptime row_layout = row_major[num_m_mmas, Int(1)]()
RowMaxTensorType
comptime RowMaxTensorType = TileTensor[dtype, Layout[*(), *()], MutUntrackedOrigin, address_space=AddressSpace.LOCAL]
RowSumTensorType
comptime RowSumTensorType = SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].RowMaxTensorType
rowwise_lanes_stride
comptime rowwise_lanes_stride = SIMD(Int(16))
score_frag_layout
comptime score_frag_layout = row_major[SoftmaxRDNA[dtype, num_m_mmas, num_n_mmas, num_warps_m, num_warps_n, use_exp2].num_colwise_tiles, Int(1)]()
ScoreFragTensorType
comptime ScoreFragTensorType = TileTensor[dtype, Layout[*(), *()], MutUntrackedOrigin, address_space=AddressSpace.LOCAL]
warp_cols
comptime warp_cols = ComptimeInt[Int(2)].static_value
warp_rows
comptime warp_rows = ComptimeInt[Int(16)].static_value
WarpLayoutT
comptime WarpLayoutT = Layout[*(), *()]
Methods
__init__
def __init__(out self)
calculate_qk_max
def calculate_qk_max(self, score: TileTensor[dtype, Storage=score.Storage, address_space=score.address_space, linear_idx_type=score.linear_idx_type], warp_scratch: TileTensor[dtype, Storage=warp_scratch.Storage, address_space=warp_scratch.address_space, linear_idx_type=warp_scratch.linear_idx_type])
calculate_qk_sum
def calculate_qk_sum(self, score: TileTensor[dtype, Storage=score.Storage, address_space=score.address_space, linear_idx_type=score.linear_idx_type], warp_scratch: TileTensor[dtype, Storage=warp_scratch.Storage, address_space=warp_scratch.address_space, linear_idx_type=warp_scratch.linear_idx_type])
exp
def exp(self, score: TileTensor[dtype, Storage=score.Storage, address_space=score.address_space, linear_idx_type=score.linear_idx_type])
calculate_correction
def calculate_correction(self)
update_output
def update_output(self, output: TileTensor[dtype, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type])
update_sum
def update_sum(self)
update_max
def update_max(self)
full
def full(self, output: TileTensor[dtype, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type], score: TileTensor[dtype, Storage=score.Storage, address_space=score.address_space, linear_idx_type=score.linear_idx_type], warp_scratch: TileTensor[dtype, Storage=warp_scratch.Storage, address_space=warp_scratch.address_space, linear_idx_type=warp_scratch.linear_idx_type])
Single-pass online softmax iteration: max -> exp -> sum -> correction -> update output -> update max/sum.
Args:
- output (
TileTensor[dtype, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type]): Accumulator tile of partial softmax outputs to correct in place with the per-row renormalization factor. - score (
TileTensor[dtype, Storage=score.Storage, address_space=score.address_space, linear_idx_type=score.linear_idx_type]): Attention score tile for this iteration; overwritten in place with exponentiated, max-subtracted values. - warp_scratch (
TileTensor[dtype, Storage=warp_scratch.Storage, address_space=warp_scratch.address_space, linear_idx_type=warp_scratch.linear_idx_type]): Shared-memory scratch tile used for cross-warp row reductions of the running maximum and sum.