Support inverse transform ue8m0 scale (#13285)
This commit is contained in:
@@ -504,6 +504,45 @@ def _transform_scale_ue8m0(sf, mn):
|
|||||||
return sf
|
return sf
|
||||||
|
|
||||||
|
|
||||||
|
def inverse_transform_scale_ue8m0(sf_packed, mn):
|
||||||
|
sf_fp32 = _inverse_transform_scale_ue8m0_impl(sf_packed)
|
||||||
|
# Can call consistency check every time since this is only called on startup
|
||||||
|
sf_packed_recreated = _transform_scale_ue8m0(sf_fp32, mn=mn)
|
||||||
|
assert torch.all(
|
||||||
|
sf_packed == sf_packed_recreated
|
||||||
|
), f"{sf_packed=} {sf_packed_recreated}"
|
||||||
|
return sf_fp32
|
||||||
|
|
||||||
|
|
||||||
|
# Inverse impl can refer to DeepGEMM's torch impl in get_mn_major_tma_aligned_packed_ue8m0_tensor_torch_impl
|
||||||
|
def _inverse_transform_scale_ue8m0_impl(sf_packed):
|
||||||
|
"""
|
||||||
|
NOTE: We assume k is aligned
|
||||||
|
:param sf_packed: (scale_mn, scale_k/4) int32
|
||||||
|
:return: (scale_mn, scale_k), float32
|
||||||
|
"""
|
||||||
|
block_size = 128
|
||||||
|
assert len(sf_packed.shape) == 2
|
||||||
|
assert sf_packed.dtype == torch.int32
|
||||||
|
|
||||||
|
mn_repeat_128, k_div_4 = sf_packed.shape
|
||||||
|
mn = mn_repeat_128 // block_size
|
||||||
|
k = k_div_4 * 4
|
||||||
|
|
||||||
|
# packed u8 -> fp32
|
||||||
|
sf_u8 = sf_packed.contiguous().flatten().view(torch.uint8).view(mn_repeat_128, k)
|
||||||
|
sf_fp32 = (sf_u8.to(torch.int32) << 23).view(torch.float32)
|
||||||
|
|
||||||
|
# remove repeat
|
||||||
|
sf_reshaped = sf_fp32.view(mn, block_size, k)
|
||||||
|
sf_unrepeated = sf_reshaped[:, 0:1, :]
|
||||||
|
assert torch.all(sf_unrepeated == sf_reshaped)
|
||||||
|
sf_unrepeated = sf_unrepeated.squeeze(1).contiguous()
|
||||||
|
|
||||||
|
assert sf_unrepeated.shape == (mn, k)
|
||||||
|
return sf_unrepeated
|
||||||
|
|
||||||
|
|
||||||
# COPIED FROM DeepGEMM
|
# COPIED FROM DeepGEMM
|
||||||
def per_block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
def per_block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
assert x.dim() == 2
|
assert x.dim() == 2
|
||||||
|
|||||||
@@ -75,6 +75,7 @@ suites = {
|
|||||||
TestFile("test_eval_fp8_accuracy.py", 303),
|
TestFile("test_eval_fp8_accuracy.py", 303),
|
||||||
TestFile("test_fa3.py", 420),
|
TestFile("test_fa3.py", 420),
|
||||||
TestFile("test_flashmla.py", 230),
|
TestFile("test_flashmla.py", 230),
|
||||||
|
TestFile("test_fp8_utils.py", 5),
|
||||||
TestFile("rotary_embedding/test_mrope.py", 10),
|
TestFile("rotary_embedding/test_mrope.py", 10),
|
||||||
TestFile("test_function_call_parser.py", 10),
|
TestFile("test_function_call_parser.py", 10),
|
||||||
TestFile("test_fused_moe.py", 80),
|
TestFile("test_fused_moe.py", 80),
|
||||||
|
|||||||
@@ -0,0 +1,44 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
|
inverse_transform_scale_ue8m0,
|
||||||
|
quant_weight_ue8m0,
|
||||||
|
transform_scale_ue8m0,
|
||||||
|
)
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class TestInverseTransformScaleUe8m0(CustomTestCase):
|
||||||
|
def test_round_trip(self):
|
||||||
|
for _ in range(100):
|
||||||
|
weight_bf16 = torch.randn(
|
||||||
|
# DeepSeek V3 kv_b_proj
|
||||||
|
(32768, 512),
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device="cuda",
|
||||||
|
)
|
||||||
|
|
||||||
|
weight_block_size = [128, 128]
|
||||||
|
|
||||||
|
qweight, sf_fp32_original = quant_weight_ue8m0(
|
||||||
|
weight_bf16, weight_block_size=weight_block_size
|
||||||
|
)
|
||||||
|
mn = qweight.shape[-2]
|
||||||
|
|
||||||
|
sf_packed_original = transform_scale_ue8m0(sf_fp32_original, mn=mn)
|
||||||
|
sf_fp32_recreated = inverse_transform_scale_ue8m0(sf_packed_original, mn=mn)
|
||||||
|
|
||||||
|
sf_packed_recreated = transform_scale_ue8m0(sf_fp32_recreated, mn=mn)
|
||||||
|
|
||||||
|
assert torch.all(
|
||||||
|
sf_packed_original == sf_packed_recreated
|
||||||
|
), f"{sf_packed_original=} {sf_packed_recreated}"
|
||||||
|
assert torch.all(
|
||||||
|
sf_fp32_original == sf_fp32_recreated
|
||||||
|
), f"{sf_fp32_original=} {sf_fp32_recreated}"
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user