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_prepared from prepare_nvfp4_w1_data(); omitting it retains the existing path. The 128- and 512-token routes also take w2_data_prepared from prepare_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 passes w2_data_prepared_k256 and w2_scale_prepared_k256 from prepare_nvfp4_w2_data_k256() / prepare_nvfp4_w2_scales_k256() next to w1_data_prepared; without them they retain the existing path. For the adjacent gate/up M512 route, pass both w1_gate_up_data_prepared and w1_gate_up_scale_prepared from 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. With accumulate=False the weighted route sum is written directly (the prior contents of out are 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] with N divisible 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_m entries.

  • 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 + 2 entries.

  • 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() and prepare_nvfp4_w2_scales().

  • w2_scale_prepared (Optional[torch.Tensor]) – Immutable prepared scale panels from prepare_nvfp4_w1_scales() and prepare_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() and prepare_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() and prepare_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() and prepare_nvfp4_w2_scales_k256().

  • w2_scale_prepared_k256 (Optional[torch.Tensor]) – Immutable K256 W2 data and scale panels from prepare_nvfp4_w2_data_k256() and prepare_nvfp4_w2_scales_k256().

  • accumulate (bool) – If True, add routed contributions to the existing out values; otherwise overwrite out with the routed sum.

Returns:

The same out tensor after routed accumulation.

Return type:

torch.Tensor