Fix B1 test LSE reference shape handling
This commit is contained in:
@@ -295,16 +295,20 @@ def test_fmha_mixed_fp8_decode():
|
||||
mag_ratio = mixed_max / ref_max if ref_max > 0 else 0.0
|
||||
|
||||
# LSE comparison
|
||||
ref_scores = torch.matmul(q_f.squeeze(2), k_f.squeeze(1).transpose(-2, -1)) * scale
|
||||
ref_lse = torch.logsumexp(ref_scores, dim=-1) # (B, H, 1)
|
||||
q_3d = q_f.squeeze(2) # (B, H, HD)
|
||||
k_3d = k_f.squeeze(1) # (B, N, HD)
|
||||
ref_scores = torch.matmul(q_3d, k_3d.transpose(-2, -1)) * scale # (B, H, N)
|
||||
ref_lse = torch.logsumexp(ref_scores, dim=-1) # (B, H)
|
||||
|
||||
passed = cos_global >= 0.999
|
||||
status = "PASS" if passed else "FAIL"
|
||||
print(f" {status}: cos_global={cos_global:.6f} min_head={min_cos:.6f} "
|
||||
f"mean_head={mean_cos:.6f}")
|
||||
print(f" |mixed|={mixed_max:.4f} |ref|={ref_max:.4f} ratio={mag_ratio:.4f}")
|
||||
print(f" LSE: mixed={lse[0,0,0].item():.4f} ref={ref_lse[0,0,0].item():.4f} "
|
||||
f"diff={abs(lse[0,0,0].item() - ref_lse[0,0,0].item()):.4f}")
|
||||
mixed_lse_val = lse.flatten()[0].item()
|
||||
ref_lse_val = ref_lse[0, 0].item()
|
||||
print(f" LSE: mixed={mixed_lse_val:.4f} ref={ref_lse_val:.4f} "
|
||||
f"diff={abs(mixed_lse_val - ref_lse_val):.4f}")
|
||||
|
||||
if not passed:
|
||||
all_pass = False
|
||||
|
||||
Reference in New Issue
Block a user