flashinfer.gemm.fp8_linear_swiglu

flashinfer.gemm.fp8_linear_swiglu(a, weight, a_scale, weight_scale, bias=None, *, config=None, out_dtype=torch.bfloat16, output_quant_scale=None, out=None)

FP8 fused linear and SwiGLU with per-token and per-channel scales.

Weight, scale, and bias rows are adjacent gate and activation pairs. The logical output width is N / 2. Supported on SM100, SM103, and SM107.

Parameters:
  • a – [M, K] torch.float8_e4m3fn activation. K must be a multiple of 128.

  • weight – [N, K] torch.float8_e4m3fn weight. N must be even.

  • a_scale – [M] float32 per-token scale.

  • weight_scale – [N] float32 per-channel scale.

  • bias – Optional [N] bfloat16 bias.

  • config – Optional PrimsTsGemmConfig matching this call. None selects the autotuned tactic.

  • out_dtype – torch.bfloat16 or torch.float8_e4m3fn. Defaults to bfloat16.

  • 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 / 2].

Returns:

The output tensor. When out_dtype is FP8 and output_quant_scale is omitted, returns (output, scale).