forked from Karylab-cklius/vllm
[ROCm] aiter_unified_attn fp8 q scale refactor (#38296)
Signed-off-by: Divakar Verma <divakar.verma@amd.com>
This commit is contained in:
@@ -200,19 +200,9 @@ class RocmAiterUnifiedAttentionImpl(RocmAttentionImpl):
|
||||
key_cache, value_cache = kv_cache.unbind(0)
|
||||
|
||||
softmax_scale = self.scale
|
||||
fp8_post_attn_v_rescale = False
|
||||
if is_quantized_kv_cache(self.kv_cache_dtype):
|
||||
key_cache = key_cache.view(self.fp8_dtype)
|
||||
value_cache = value_cache.view(self.fp8_dtype)
|
||||
# When Q is FP8, triton kernel skips K/V dequant (for fp8xfp8 matmul).
|
||||
# Compensate by absorbing q_scale and k_scale into softmax_scale, and
|
||||
# v_scale into output_scale (or post-multiplying if no fusion).
|
||||
if query.dtype == self.fp8_dtype:
|
||||
softmax_scale = self.scale * layer._q_scale_float * layer._k_scale_float
|
||||
if output_scale is not None:
|
||||
output_scale = output_scale / layer._v_scale_float
|
||||
else:
|
||||
fp8_post_attn_v_rescale = True
|
||||
|
||||
cu_seqlens_q = attn_metadata.query_start_loc
|
||||
seqused_k = attn_metadata.seq_lens
|
||||
@@ -220,11 +210,6 @@ class RocmAiterUnifiedAttentionImpl(RocmAttentionImpl):
|
||||
max_seqlen_k = attn_metadata.max_seq_len
|
||||
block_table = attn_metadata.block_table
|
||||
|
||||
descale_shape = (
|
||||
cu_seqlens_q.shape[0] - 1,
|
||||
key.shape[1] if key is not None else self.num_kv_heads,
|
||||
)
|
||||
|
||||
self.unified_attention(
|
||||
q=query[:num_actual_tokens],
|
||||
k=key_cache,
|
||||
@@ -240,16 +225,13 @@ class RocmAiterUnifiedAttentionImpl(RocmAttentionImpl):
|
||||
window_size=self.sliding_window,
|
||||
block_table=block_table,
|
||||
softcap=self.logits_soft_cap,
|
||||
q_descale=None, # q_scale absorbed into softmax_scale
|
||||
k_descale=layer._k_scale.expand(descale_shape),
|
||||
v_descale=layer._v_scale.expand(descale_shape),
|
||||
q_descale=layer._q_scale if query.dtype == self.fp8_dtype else None,
|
||||
k_descale=layer._k_scale,
|
||||
v_descale=layer._v_scale,
|
||||
sinks=self.sinks,
|
||||
output_scale=output_scale,
|
||||
)
|
||||
|
||||
if fp8_post_attn_v_rescale:
|
||||
output[:num_actual_tokens].mul_(layer._v_scale_float)
|
||||
|
||||
return output
|
||||
|
||||
def do_kv_cache_update(
|
||||
|
||||
Reference in New Issue
Block a user