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

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. .uint8 for FP4).
  • ​scales_dtype (DType): Scale factor dtype (FP8 variant).
  • ​output_layout (TensorLayout): Layout of the quantized output TileTensor.
  • ​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 selects get_device_alignment().
  • ​_n_warps (Int): Number of warps that concurrently stage tiles through the dispatch scatter SMEM. Sizes dispatch_smem_size and 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 to EP_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 at sf_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 as dst_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-wide input_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 in output_rowwise_scales. A token then dequantizes as fp4 * 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:

get_group_size​

static def get_group_size() -> Int

Returns:

Int

get_type_name​

static def get_type_name() -> String

Returns:

String

quant_size​

static def quant_size() -> Int

Returns:

Int

scales_size​

static def scales_size() -> Int

Returns:

Int

token_size​

static def token_size() -> Int

Returns:

Int

scales_offset​

static def scales_offset() -> Int

Returns:

Int

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:

Int

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:

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)

Was this page helpful?