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 K must 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.bfloat16 or torch.float8_e4m3fn. Defaults to bfloat16. NVFP4 output is only available from fp4_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 uses 1 and 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_dtype is FP8 and output_quant_scale is omitted, returns (output, scale).