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_e4m3fnactivation.Kmust be a multiple of 128.weight –
[N, K]torch.float8_e4m3fnweight.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 whenNis split into three equal groups.config – Optional
PrimsTsGemmConfigmatching this call.Noneselects the autotuned tactic.out_dtype –
torch.bfloat16ortorch.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 uses1and 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_dtypeis FP8 andoutput_quant_scaleis omitted, returns(output, scale).