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. N must 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, or torch.uint8 for packed NVFP4 output. Defaults to bfloat16. NVFP4 output requires logical N / 2 to 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 uses 1 and 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_scale returns (output, scale). NVFP4 output returns (output, block_scale), or (output, block_scale, scale) when output_quant_scale is omitted.