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
NVBlockScaledTokenFormat
struct NVBlockScaledTokenFormat[quant_dtype: DType, scales_dtype: DType, output_layout: TensorLayout, scales_offset_layout: TensorLayout, //, _hid_dim: Int, _top_k: Int, _alignment: Int = Int(0), _n_warps: Int = Int(32), ep_copy_role_split: Bool = False, copy_src_width: Int = Int(8), ep_prequantized: Bool = False, sf_arena_rows: Int = Int(0), sf_vec4: Bool = False, sf_kxg: Int = Int(0), sf_rows_per_expert: Int = Int(0), nvfp4_dyn_global_scales: Bool = False]
Token format for NVIDIA block-scaled FP4/FP8 quantization.
Supports NVFP4 (quant_dtype=uint8, scales_dtype=float8_e4m3fn),
MXFP4 (quant_dtype=uint8, scales_dtype=float8_e8m0fnu), and MXFP8
(quant_dtype=float8_e4m3fn, scales_dtype=float8_e8m0fnu) wire formats.
Uses TMA-based copies for scale preshuffle into the output tensor.
Parameters
- quant_dtype (
DType): Quantized element dtype (e.g..uint8for FP4). - scales_dtype (
DType): Scale factor dtype (FP8 variant). - output_layout (
TensorLayout): Layout of the quantized outputTileTensor. - scales_offset_layout (
TensorLayout): Layout of the per-expert scale offset tensor. - _hid_dim (
Int): Hidden dimension; must be divisible by the group size. - _top_k (
Int): Number of experts each token is routed to. - _alignment (
Int): Override for the byte alignment of the wire buffer; 0 selectsget_device_alignment(). - _n_warps (
Int): Number of warps that concurrently stage tiles through the dispatch scatter SMEM. Sizesdispatch_smem_sizeand the mbar / tile-buffer layout. Defaults to 32 (the standalone dispatch block of 1024 threads); the co-resident megakernel sets it to its comm warp count so the staging does not oversize the FFN pipeline reservation. - ep_copy_role_split (
Bool): Whether the per-token copy body uses the split publisher fan-out. Must match the dispatch kernel's parameter of the same name. Off by default, which keeps the stock warp-strided fan-out. - copy_src_width (
Int): Elements moved per source step by the copy body. Defaults toEP_COPY_SRC_WIDTH; must divide the scale group size. - ep_prequantized (
Bool): Whether the source rows already carry canonical quantized payload and scales, so the copy moves bytes instead of quantizing. When set, the BF16 conversion, absolute-max reduction, scale calculation and UE8M0 encoding are not instantiated at all. - sf_arena_rows (
Int): Row capacity of the destination's message region. The arena starts atsf_arena_rows * msg_size()inside the peer receive allocation. Zero disables the arena scatter entirely, which keeps every other specialization bit-unchanged. - sf_vec4 (
Bool): Whether to move each scale-factor atom with one 4-byte store instead of four 1-byte stores. Requires a 4-byte atom (SF_ATOM_K == 4), which is asserted at the store site. - sf_kxg (
Int): Number of scale-factor k-tiles per row. The tiles are spread over the warp's lanes so each scale is written exactly once, and it sets the per-128-row block stride. - sf_rows_per_expert (
Int): Rows reserved per expert in the final layout; forms the arena row asdst_expert_local_idx * sf_rows_per_expert + final_row. - nvfp4_dyn_global_scales (
Bool): Whether each NVFP4 token is quantized against its own global scale,2688 / rowmax, instead of the kernel-wideinput_scale. The sender reduces the token's absolute max in a first pass and quantizes in a second; the scale's BF16 inverse travels behind the per-block scales in the message, and the receiver stores it inoutput_rowwise_scales. A token then dequantizes asfp4 * block_scale * rowwise_scale.
Fields
- scales_tma_op (
NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales].ScalesTMATensorTileType): - output_tokens (
NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales].TensorType): - output_scales_offset (
NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales].ScalesOffsetTensorType): - output_rowwise_scales (
NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales].RowwiseScalesTensorType):
Implemented traits
AnyType,
Copyable,
Deinitable,
DevicePassable,
ImplicitlyCopyable,
Movable,
TokenFormat
comptime members
alignment
comptime alignment = _alignment if _alignment.__bool__() else get_device_alignment()
device_type
comptime device_type = NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales]
dispatch_smem_size
comptime dispatch_smem_size = (Int((add (mul align_up(Coord[Int(4), DType.int64](Index[Int, Int, Int, Int](Int(1), (((_hid_dim // NVBlockScaledTokenFormat.get_group_size()) // Int(4)) // (load_from_mem Tuple(Int(128), Int(2)).__getitem_param__[Int(1)]())), Int(1), Int((mul (load_from_mem Tuple(Int(32), Int(4)).__getitem_param__[Int(1)]()), 4)))).product(), Int(128)), size_of[scales_dtype](), _n_warps), (mul align_up((NVBlockScaledTokenFormat.quant_size() // (load_from_mem Tuple(Int(128), Int(2)).__getitem_param__[Int(1)]())), Int(16)), _n_warps))) + align_up(Int((mul size_of[SharedMemBarrier](), _n_warps)), Int(8)))
dispatch_wait_tile_shape
comptime dispatch_wait_tile_shape = Tuple(Int(128), Int(2))
group_size
comptime group_size = NVBlockScaledTokenFormat.get_group_size()
hid_dim
comptime hid_dim = _hid_dim
is_mxfp4
comptime is_mxfp4 = (quant_dtype == DType.uint8) if (scales_dtype == DType.float8_e8m0fnu) else (scales_dtype == DType.float8_e8m0fnu)
is_mxfp8
comptime is_mxfp8 = (quant_dtype == DType.float8_e4m3fn) if (scales_dtype == DType.float8_e8m0fnu) else (scales_dtype == DType.float8_e8m0fnu)
is_nvfp4
comptime is_nvfp4 = (quant_dtype == DType.uint8) if (scales_dtype == DType.float8_e4m3fn) else (scales_dtype == DType.float8_e4m3fn)
RowwiseScalesTensorType
comptime RowwiseScalesTensorType = TileTensor[.bfloat16, Layout[TypeList[Int64](), TypeList[ComptimeInt[Int(1)]]()], MutUntrackedOrigin]
ScalesOffsetTensorType
comptime ScalesOffsetTensorType = TileTensor[.uint32, scales_offset_layout, MutUntrackedOrigin]
ScalesTMATensorTileType
comptime ScalesTMATensorTileType = TMATensorTile[scales_dtype, Int(4), NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales].tma_tile_shape, _default_desc_shape[Int(4), scales_dtype, NVBlockScaledTokenFormat[_hid_dim, _top_k, _alignment, _n_warps, ep_copy_role_split, copy_src_width, ep_prequantized, sf_arena_rows, sf_vec4, sf_kxg, sf_rows_per_expert, nvfp4_dyn_global_scales].tma_tile_shape, TensorMapSwizzle.SWIZZLE_NONE]()]
TensorType
comptime TensorType = TileTensor[quant_dtype, output_layout, MutUntrackedOrigin]
tma_tile_shape
comptime tma_tile_shape = Index[Int, Int, Int, Int](Int(1), (((_hid_dim // NVBlockScaledTokenFormat.get_group_size()) // Int(4)) // (load_from_mem Tuple(Int(128), Int(2)).__getitem_param__[Int(1)]())), Int(1), (Int(4) * (load_from_mem SF_ATOM_M.__getitem_param__[Int(1)]())))
top_k
comptime top_k = _top_k
Methods
__init__
def __init__(out self, output_tokens: TileTensor[quant_dtype, output_layout, address_space=output_tokens.address_space, linear_idx_type=output_tokens.linear_idx_type], output_scales: TileTensor[scales_dtype, address_space=output_scales.address_space, linear_idx_type=output_scales.linear_idx_type], output_scales_offset: TileTensor[.uint32, scales_offset_layout, address_space=output_scales_offset.address_space, linear_idx_type=output_scales_offset.linear_idx_type], ctx: DeviceContext, output_rowwise_scales: Optional[Pointer[BFloat16, MutAnyOrigin]] = None)
Wraps the dispatch outputs and builds the scales TMA descriptor.
Args:
- output_tokens (
TileTensor[quant_dtype, output_layout, address_space=output_tokens.address_space, linear_idx_type=output_tokens.linear_idx_type]): Quantized tokens, one row per received token. - output_scales (
TileTensor[scales_dtype, address_space=output_scales.address_space, linear_idx_type=output_scales.linear_idx_type]): Per-block scale tiles in the grouped matmul's 5D scale-factor layout. - output_scales_offset (
TileTensor[.uint32, scales_offset_layout, address_space=output_scales_offset.address_space, linear_idx_type=output_scales_offset.linear_idx_type]): Per-expert offsets into the scale tiles. - ctx (
DeviceContext): Device context the TMA descriptor is created on. - output_rowwise_scales (
Optional[Pointer[BFloat16, MutAnyOrigin]]): Per-token inverse global scales, one BF16 per row ofoutput_tokens. Required withnvfp4_dyn_global_scalesand ignored otherwise.
get_group_size
get_type_name
quant_size
scales_size
token_size
scales_offset
global_scale_offset
static def global_scale_offset() -> Int
Returns the message offset of the token's BF16 inverse global scale, which sits directly behind the per-block scales.
Returns:
pad_expert_offsets
def pad_expert_offsets[n_groups: Int](self, row_offsets: Pointer[UInt32, address_space=row_offsets.address_space])
The mojo NVFP4 grouped matmul doesn't require padding for each group's FP4 quants. However, it requires each group's scales to be aligned to the SF_MN_GROUP_SIZE=128. This function updates the output_scales_offset tensor to satisfy this requirement.
For example, if the row_offsets tensor is [0, 100, 300, 400], this function will update the output_scales_offset tensor to [0, 1, 1]. The formula is: For group i, its first scales block index is row_offsets[i] // SF_MN_GROUP_SIZE + output_scales_offset[i]. Group 0, 1 and 2 have 100, 200, 100 tokens respectively, so the number of scales blocks are 1, 2, 1 respectively. The scales blocks for group 1 start at 100 // 128 + output_scales_offset[1] = 1, and the scales blocks for group 2 start at 300 // 128 + output_scales_offset[2] = 3.
copy_token_to_send_buf
static def copy_token_to_send_buf[src_type: DType, block_size: Int, buf_addr_space: AddressSpace = .GENERIC, thread_base: Int = Int(0)](buf_p: Pointer[UInt8, address_space=buf_addr_space], src_p: Pointer[Scalar[src_type], address_space=src_p.address_space], input_scale: Float32)
scatter_row_scales
static def scatter_row_scales(recv_base: Pointer[UInt8, MutUntrackedOrigin], send_buf_p: Pointer[UInt8, MutUntrackedOrigin], dst_expert_local_idx: Int32, final_row: Int32, lane: Int)
Scatters this row's scales into the SF-atom arena.
Same warp, same final row, no extra reservation and no per-row
atomic: the address is a pure function of the row the caller already
reserved. One store per (row, k-tile); sf_kxg k-tiles are spread
over the warp's lanes, so every scale value is written exactly once.
The source is this format's own scale region inside the staged
message (scales_offset()), and the destination arena begins after
the message region, whose extent is sf_arena_rows * msg_size() --
both quantities the format already defines, so the caller supplies
only the peer base and the routing it already computed.
Args:
- recv_base (
Pointer[UInt8, MutUntrackedOrigin]): Base of the destination peer's receive allocation. - send_buf_p (
Pointer[UInt8, MutUntrackedOrigin]): This row's staged send-buffer message. - dst_expert_local_idx (
Int32): Destination-local expert index. - final_row (
Int32): Reserved row in the destination's final layout. - lane (
Int): Calling lane; spreads the k-tiles across the warp.
copy_msg_to_output_tensor
def copy_msg_to_output_tensor[buf_addr_space: AddressSpace = .GENERIC](self, buf_p: Pointer[UInt8, address_space=buf_addr_space], token_index: Int, expert_slot: Int = Int(0), expert_start: Int = Int(0))
NVFP4 format directly uses tile based copy.
init_smem_resources
def init_smem_resources[smem_base_offset: Int = Int(0), warp_base: Int = Int(0)](self)
copy_msg_tile_to_output_tensor
def copy_msg_tile_to_output_tensor[extract_topk_info_func: def(Pointer[UInt8, MutUntrackedOrigin], Int) -> None, recv_buf_ptr_func: def(Int) -> Pointer[UInt8, MutUntrackedOrigin], //, n_warps: Int, shared_expert_offset: Int = Int(0), warp_base: Int = Int(0), smem_base_offset: Int = Int(0)](self, expert_id: Int, expert_start_pos: Int, tile_id: Int, tile_end: Int, extract_topk_info_functor: extract_topk_info_func, recv_buf_ptr_functor: recv_buf_ptr_func)