[DSv4] Use BF16 instead of FP32 for indexer score computation (#30012)

Co-authored-by: 百麒 <yaozhong.lyz@alibaba-inc.com>
This commit is contained in:
Lewis
2026-07-15 15:06:55 -07:00
committed by GitHub
co-authored by 百麒
parent 26cb0fcdda
commit 3101c1258c
@@ -87,14 +87,14 @@ def fp8_paged_mqa_logits_torch(
kv_values_raw = kvcache_gathered[..., :SCALE_OFFSET].contiguous() kv_values_raw = kvcache_gathered[..., :SCALE_OFFSET].contiguous()
kv_values_fp8 = kv_values_raw.view(dtype=FP8_DTYPE) kv_values_fp8 = kv_values_raw.view(dtype=FP8_DTYPE)
kv_values = kv_values_fp8.to(torch.float32) kv_values = kv_values_fp8.to(torch.bfloat16)
kv_values = kv_values.reshape(batch_size, max_num_pages * block_size, head_dim) kv_values = kv_values.reshape(batch_size, max_num_pages * block_size, head_dim)
kv_scales_raw = kvcache_gathered[..., SCALE_OFFSET:].contiguous() kv_scales_raw = kvcache_gathered[..., SCALE_OFFSET:].contiguous()
kv_scales = kv_scales_raw.view(dtype=torch.float32) kv_scales = kv_scales_raw.view(dtype=torch.float32)
kv_scales = kv_scales.reshape(batch_size, max_num_pages * block_size) kv_scales = kv_scales.reshape(batch_size, max_num_pages * block_size)
q_float = q_fp8[:, 0].to(torch.float32) q_float = q_fp8[:, 0].to(torch.bfloat16)
scores = torch.bmm(kv_values, q_float.transpose(1, 2)) scores = torch.bmm(kv_values, q_float.transpose(1, 2))
scores = F.relu(scores) scores = F.relu(scores)
scores = scores * weight.unsqueeze(1) scores = scores * weight.unsqueeze(1)
@@ -199,13 +199,13 @@ def fp8_paged_mqa_logits_torch_sm120(
kv_value_raw = kvcache_gathered[..., :SCALE_OFFSET] kv_value_raw = kvcache_gathered[..., :SCALE_OFFSET]
kv_scale_raw = kvcache_gathered[..., SCALE_OFFSET:] kv_scale_raw = kvcache_gathered[..., SCALE_OFFSET:]
kv_value = kv_value_raw.contiguous().view(dtype=FP8_DTYPE).to(torch.float32) kv_value = kv_value_raw.contiguous().view(dtype=FP8_DTYPE).to(torch.bfloat16)
kv_value = kv_value.view(batch_size, max_padded_seq, head_dim) kv_value = kv_value.view(batch_size, max_padded_seq, head_dim)
kv_scale = kv_scale_raw.contiguous().view(dtype=torch.float32) kv_scale = kv_scale_raw.contiguous().view(dtype=torch.float32)
kv_scale = kv_scale.view(batch_size, max_padded_seq) kv_scale = kv_scale.view(batch_size, max_padded_seq)
q = q_fp8[:, 0].to(torch.float32) q = q_fp8[:, 0].to(torch.bfloat16)
score = torch.bmm(kv_value, q.transpose(1, 2)) score = torch.bmm(kv_value, q.transpose(1, 2))