fix test_weight_checker_comparator assertion and ue8m0 scale unpack (#29623)
This commit is contained in:
@@ -59,6 +59,8 @@ class Fp8BlockComparable(ComparableWeight):
|
|||||||
def _normalize_scale(w_q: torch.Tensor, w_s: torch.Tensor) -> torch.Tensor:
|
def _normalize_scale(w_q: torch.Tensor, w_s: torch.Tensor) -> torch.Tensor:
|
||||||
if w_s.dtype == torch.int32:
|
if w_s.dtype == torch.int32:
|
||||||
w_s = inverse_transform_scale_ue8m0(w_s, mn=w_q.shape[-2])
|
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)
|
return w_s.to(torch.float32)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -133,7 +133,10 @@ class TestCompareQuantPair(CustomTestCase):
|
|||||||
reference = _compare_quant_pair(self.e_q, self.e_s, self.a_q, self.a_s)
|
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):
|
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)
|
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
|
@staticmethod
|
||||||
def _quantize_partial(weight: torch.Tensor, scale_margin: float):
|
def _quantize_partial(weight: torch.Tensor, scale_margin: float):
|
||||||
|
|||||||
Reference in New Issue
Block a user