flashinfer.diffusion_ops.prepare_minimax_h3_fc1_weight_fp8

flashinfer.diffusion_ops.prepare_minimax_h3_fc1_weight_fp8(fc1_weight: Tensor, chunk_rows: int = 2048) → Tuple[Tensor, Tensor]

Quantize the fused FC1 weight for minimax_h3_fc1_swiglu_fp8() (SM120).

fc1_weight is the BF16 [28672, 5376] matrix with gate rows [0, 14336) followed by up rows [14336, 28672). Each output channel (row) is quantized to E4M3 with scale = RN(max(amax(row), 1e-12) / 448) (true IEEE division) and q = RN_sat(row / scale), then the rows are permuted into the SM120 prepacked order (interleave_minimax_h3_fc1_rows_sm120(): eight gate rows, then the eight up rows of the same output columns).

Returns (fc1_weight_q, fc1_weight_scale): float8_e4m3fn [28672, 5376] and float32 [28672] in the prepacked row order. The layout is specific to the SM120 operator; it is not interchangeable with the SM100/SM103 prepare_minimax_h3_fc1_weight_* outputs.

Parameters:
  • fc1_weight (torch.Tensor) – BF16 CUDA tensor with shape [28672, 5376] in gate-rows-then-up-rows order.

  • chunk_rows (int) – Number of weight rows quantized per temporary FP32 chunk.

Returns:

The prepacked E4M3 weight and its FP32 per-row scales.

Return type:

Tuple[torch.Tensor, torch.Tensor]