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_weightis the BF16[28672, 5376]matrix with gate rows[0, 14336)followed by up rows[14336, 28672). Each output channel (row) is quantized to E4M3 withscale = RN(max(amax(row), 1e-12) / 448)(true IEEE division) andq = 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]andfloat32[28672]in the prepacked row order. The layout is specific to the SM120 operator; it is not interchangeable with the SM100/SM103prepare_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]