flashinfer.gemm.fp8_linear

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

FP8 [M, K] times [N, K] with per-token and per-channel scales.

Supported on SM100, SM103, and SM107. There is no global scale.

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

  • weight – [N, K] torch.float8_e4m3fn weight.

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

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

  • bias – Optional [N] bfloat16 bias.

  • qkv_scale – Optional [3] float32 scale applied per Q, K, and V group when N is split into three equal groups.

  • 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. BF16 shape is [M, N]. FP8 shape is [M, N].

Returns:

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