flashinfer.gemm.fp8_qkv_qknorm_rope¶
- flashinfer.gemm.fp8_qkv_qknorm_rope(a, qkv_weight, a_scale, qkv_weight_scale, q_norm_weight, k_norm_weight, cos_sin, positions, *, num_q_heads, num_kv_heads, head_dim, is_neox=False, qkv_scale=None, config=None, out_dtype=torch.bfloat16, output_quant_scale=None, out=None)¶
FP8 QKV projection with QK RMSNorm and RoPE.
positions[i]selects the row ofcos_sinused for tokeni. Values outside[0, M)use row 0. Supported on SM100, SM103, and SM107.num_q_headsmust equalnum_kv_heads,head_dimmust be 128, andis_neoxmust be false.- Parameters:
a –
[M, K]torch.float8_e4m3fnactivation.Kmust be a multiple of 128.qkv_weight –
[N, K]torch.float8_e4m3fnweight packed as three equal Q, K, and V groups.Nis3 * num_q_heads * head_dim.a_scale –
[M]float32 per-token scale.qkv_weight_scale –
[N]float32 per-channel scale.q_norm_weight –
[head_dim]bfloat16 RMSNorm weight for Q.k_norm_weight –
[head_dim]bfloat16 RMSNorm weight for K.cos_sin –
[M, head_dim]float32 table. The first half of each row is cosine and the second half is sine.positions –
[M]int64 indices intocos_sin.num_q_heads – Number of query heads. Must equal
num_kv_heads.num_kv_heads – Number of key and value heads.
head_dim – Head size. Only 128 is supported.
is_neox – RoPE layout. Only false, the interleaved-pair layout, is supported.
qkv_scale – Optional
[3]float32 scale for the Q, K, and V 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 of shape
[M, N].
- Returns:
The output tensor. When
out_dtypeis FP8 andoutput_quant_scaleis omitted, returns(output, scale).