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 of cos_sin used for token i. Values outside [0, M) use row 0. Supported on SM100, SM103, and SM107. num_q_heads must equal num_kv_heads, head_dim must be 128, and is_neox must be false.

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

  • qkv_weight – [N, K] torch.float8_e4m3fn weight packed as three equal Q, K, and V groups. N is 3 * 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 into cos_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 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].

Returns:

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