flashinfer.gemm.fp4_linear_swiglu¶
- flashinfer.gemm.fp4_linear_swiglu(a_packed, a_block_scale, a_global_scale, weight_packed, weight_block_scale, weight_global_scale, bias=None, *, out_dtype=torch.bfloat16, output_quant_scale=None, out=None)¶
Packed NVFP4 fused linear and SwiGLU.
Block scales are contiguous 1D 128x4 FP8-E4M3 buffers. Weight rows are adjacent gate and activation pairs, and the logical output width is
N / 2. Supported on SM100, SM103, and SM107.- Parameters:
a_packed –
[M, K / 2]uint8 activation, two FP4 values per byte.a_block_scale – 128x4 block scales for
a_packed.a_global_scale – Host float multiplied into the activation scale.
weight_packed –
[N, K / 2]uint8 weight.Nmust be even.weight_block_scale – 128x4 block scales for
weight_packed.weight_global_scale – Host float multiplied into the weight scale.
bias – Optional
[N]bfloat16 bias.out_dtype –
torch.bfloat16,torch.float8_e4m3fn, ortorch.uint8for packed NVFP4 output. Defaults to bfloat16. NVFP4 output requires logicalN / 2to be a multiple of 16.output_quant_scale – Optional
[1]float32 encode scale for FP8 or NVFP4 output. A supplied tensor is read on every call. When omitted, the kernel uses1and the returned scale is a separate tensor; writing it does not change later calls.out – Optional preallocated output. BF16 and FP8 shape is
[M, N / 2]. Packed NVFP4 shape is[M, N / 4].
- Returns:
The output tensor. FP8 output with no
output_quant_scalereturns(output, scale). NVFP4 output returns(output, block_scale), or(output, block_scale, scale)whenoutput_quant_scaleis omitted.