flashinfer.fused_moe.alphamoe_nvfp4_routed_moe¶
- flashinfer.fused_moe.alphamoe_nvfp4_routed_moe(hidden_states: Tensor, hidden_states_scale: Tensor, gemm1_weights: Tensor, gemm1_weights_scale: Tensor, gemm2_weights: Tensor, gemm2_weights_scale: Tensor, output1_scale_gate_scalar: Tensor, output1_scale_scalar: Tensor, output2_scale_scalar: Tensor, sorted_token_ids: Tensor, expert_ids: Tensor, num_tokens_post_padded: Tensor, topk_weights: Tensor, out: Tensor, topk_ids: Tensor, cumsum_buffer: Tensor, top_k: int, block_m: int = 8, routed_scaling_factor: float = 1.0, w1_scale_prepared: Tensor | None = None, w2_scale_prepared: Tensor | None = None, w1_data_prepared: Tensor | None = None, w1_gate_up_data_prepared: Tensor | None = None, w1_gate_up_scale_prepared: Tensor | None = None, w1_scale_prepared_interleaved: Tensor | None = None, w2_data_prepared: Tensor | None = None, w2_data_prepared_k256: Tensor | None = None, w2_scale_prepared_k256: Tensor | None = None, accumulate: bool = True) Tensor¶
Run route alignment, complete expert computation and weighted accumulation.
Selected complete routes use private per-route accumulators and finalize contributions in route order, preserving the caller’s existing output. The exact route controls seed initialization, FP32 or BF16 route storage, optional compact-owner metadata and immutable prepared W1 scales/data. Buffer initialization, expert computation and finalization all stay inside this call and use the current stream, including during CUDA graph capture.
Inputs without a selected complete route retain the existing alignment, accumulator and compute path. The aligned API is unchanged. Prepared scale tensors remain optional immutable model-load inputs; this function performs no weight preparation, host device-count read or automatic model fetching. For the selected 8-, 128- and 512-token routes, pass
w1_data_preparedfromprepare_nvfp4_w1_data(); omitting it retains the existing path. The 128- and 512-token routes also takew2_data_preparedfromprepare_nvfp4_w2_data(); without both prepared data tensors those token counts retain the existing path. Token counts above 512 take the token-tile route when the caller also passesw2_data_prepared_k256andw2_scale_prepared_k256fromprepare_nvfp4_w2_data_k256()/prepare_nvfp4_w2_scales_k256()next tow1_data_prepared; without them they retain the existing path. For the adjacent gate/up M512 route, pass bothw1_gate_up_data_preparedandw1_gate_up_scale_preparedfrom the matching model-load preparation helpers. They use a distinct layout from the optional old prepared inputs. No packing occurs in this call. The caller-owned output is updated and returned. Withaccumulate=Falsethe weighted route sum is written directly (the prior contents ofoutare ignored, so no zero fill is needed); the result equals the seeded path on a zero output bit for bit.- Parameters:
hidden_states (torch.Tensor) – Packed E2M1 activations
[M, K / 2]. The innermost stride must be 1; row-strided views are supported.hidden_states_scale (torch.Tensor) – Linear E4M3 scales
[M, K / 16].gemm1_weights (torch.Tensor) – Packed gate/up weights
[E, N, K / 2]withNdivisible by 256.gemm1_weights_scale (torch.Tensor) – Linear E4M3 scales
[E, N, K / 16].gemm2_weights (torch.Tensor) – Packed down weights
[E, K, N / 4].gemm2_weights_scale (torch.Tensor) – Linear E4M3 scales
[E, K, N / 32].output1_scale_gate_scalar (torch.Tensor) – Contiguous FP32 per-expert gate dequantization scales
[E].output1_scale_scalar (torch.Tensor) – Contiguous FP32 per-expert up-projection scales
[E].output2_scale_scalar (torch.Tensor) – Contiguous FP32 per-expert down-projection scales
[E].sorted_token_ids (torch.Tensor) – Caller-owned contiguous int32 storage for aligned-plan entries. It must hold at least
expert_ids.numel() * block_mentries.expert_ids (torch.Tensor) – Caller-owned contiguous int32 storage for expert ids in the aligned plan.
num_tokens_post_padded (torch.Tensor) – Caller-owned one-element int32 tensor receiving the valid plan extent.
topk_weights (torch.Tensor) – FP32 route weights
[M, top_k].out (torch.Tensor) – Contiguous BF16 output
[M, K]. It is updated in place and returned.topk_ids (torch.Tensor) – Contiguous int32 routed expert ids
[M, top_k]used to build the plan.cumsum_buffer (torch.Tensor) – Caller-owned contiguous int32 routing workspace with at least
E + 2entries.top_k (int) – Routes per token.
block_m (int) – Routing-plan block size, at least 8 and divisible by 8.
routed_scaling_factor (float) – Finite scalar applied to each routed contribution.
w1_scale_prepared (Optional[torch.Tensor]) – Immutable prepared scale panels from
prepare_nvfp4_w1_scales()andprepare_nvfp4_w2_scales().w2_scale_prepared (Optional[torch.Tensor]) – Immutable prepared scale panels from
prepare_nvfp4_w1_scales()andprepare_nvfp4_w2_scales().w1_data_prepared (Optional[torch.Tensor]) – Immutable W1 panels from
prepare_nvfp4_w1_data()for eligible routes.w1_gate_up_data_prepared (Optional[torch.Tensor]) – Immutable adjacent gate/up panels from
prepare_nvfp4_w1_gate_up_data()andprepare_nvfp4_w1_gate_up_scales()for the M512 route.w1_gate_up_scale_prepared (Optional[torch.Tensor]) – Immutable adjacent gate/up panels from
prepare_nvfp4_w1_gate_up_data()andprepare_nvfp4_w1_gate_up_scales()for the M512 route.w1_scale_prepared_interleaved (Optional[torch.Tensor]) – Immutable interleaved W1 scales from
prepare_nvfp4_w1_scales_interleaved().w2_data_prepared (Optional[torch.Tensor]) – Immutable W2 panels from
prepare_nvfp4_w2_data()for eligible routes.w2_data_prepared_k256 (Optional[torch.Tensor]) – Immutable K256 W2 data and scale panels from
prepare_nvfp4_w2_data_k256()andprepare_nvfp4_w2_scales_k256().w2_scale_prepared_k256 (Optional[torch.Tensor]) – Immutable K256 W2 data and scale panels from
prepare_nvfp4_w2_data_k256()andprepare_nvfp4_w2_scales_k256().accumulate (bool) – If
True, add routed contributions to the existingoutvalues; otherwise overwriteoutwith the routed sum.
- Returns:
The same
outtensor after routed accumulation.- Return type:
torch.Tensor