flashinfer.gemm.fp4_linear¶
- flashinfer.gemm.fp4_linear(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
[M, K]times[N, K].Block scales are contiguous 1D 128x4 FP8-E4M3 buffers. Global scales are host floats. Supported on SM100, SM103, and SM107. Logical
Kmust be a multiple of 256.- 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.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.bfloat16ortorch.float8_e4m3fn. Defaults to bfloat16. NVFP4 output is only available fromfp4_linear_swiglu().output_quant_scale – Optional
[1]float32 encode scale for FP8 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 of shape
[M, N].
- Returns:
The output tensor. When
out_dtypeis FP8 andoutput_quant_scaleis omitted, returns(output, scale).