fix test_weight_checker_comparator assertion and ue8m0 scale unpack (#29623)

This commit is contained in:
Yueming Yuan
2026-06-28 23:36:37 -07:00
committed by GitHub
parent f85cc94d82
commit cd91dd0757
2 changed files with 6 additions and 1 deletions
@@ -59,6 +59,8 @@ class Fp8BlockComparable(ComparableWeight):
def _normalize_scale(w_q: torch.Tensor, w_s: torch.Tensor) -> torch.Tensor:
if w_s.dtype == torch.int32:
w_s = inverse_transform_scale_ue8m0(w_s, mn=w_q.shape[-2])
# ue8m0 packing aligns k to a multiple of 4; drop the padding blocks.
w_s = w_s[..., : -(-w_q.shape[-1] // 128)]
return w_s.to(torch.float32)
@staticmethod
@@ -133,7 +133,10 @@ class TestCompareQuantPair(CustomTestCase):
reference = _compare_quant_pair(self.e_q, self.e_s, self.a_q, self.a_s)
with patch("sglang.srt.utils.weight_checker_comparator.CHUNK_NUMEL", 128 * 128):
chunked = _compare_quant_pair(self.e_q, self.e_s, self.a_q, self.a_s)
self.assertEqual(chunked, reference)
eq_c, max_c, mean_c, ex_c = chunked
eq_r, max_r, mean_r, ex_r = reference
self.assertEqual((eq_c, max_c, ex_c), (eq_r, max_r, ex_r))
self.assertAlmostEqual(mean_c, mean_r, places=7)
@staticmethod
def _quantize_partial(weight: torch.Tensor, scale_margin: float):