Fix ModelOpt NVFP4 scalar scales for merged linears (#29151)

This commit is contained in:
Po-Han Huang (NVIDIA)
2026-07-13 16:14:20 -07:00
committed by GitHub
parent a909077d22
commit cfc3d0555e
2 changed files with 121 additions and 13 deletions
@@ -0,0 +1,80 @@
import unittest
import torch
from sglang.srt.layers.linear import MergedColumnParallelLinear, QKVParallelLinear
from sglang.srt.layers.parameter import PerTensorScaleParameter
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestModelOptNvfp4(CustomTestCase):
def _make_layer(self):
return MergedColumnParallelLinear(
input_size=16,
output_sizes=[16, 16],
bias=False,
tp_rank=0,
tp_size=1,
)
def _make_qkv_layer(self):
return QKVParallelLinear(
hidden_size=16,
head_size=8,
total_num_heads=2,
total_num_kv_heads=2,
bias=False,
tp_rank=0,
tp_size=1,
)
def test_fused_scalar_scale_load_fills_all_logical_slots(self):
layer = self._make_layer()
scale = PerTensorScaleParameter(
data=torch.empty(2, dtype=torch.float32),
weight_loader=layer.weight_loader_v2,
)
layer.weight_loader_v2(scale, torch.tensor(0.25, dtype=torch.float32))
torch.testing.assert_close(scale, torch.tensor([0.25, 0.25]))
def test_fused_scalar_scale_load_rejects_non_scalar(self):
layer = self._make_layer()
scale = PerTensorScaleParameter(
data=torch.empty(2, dtype=torch.float32),
weight_loader=layer.weight_loader_v2,
)
with self.assertRaisesRegex(ValueError, "Expected scalar scale"):
layer.weight_loader_v2(scale, torch.tensor([0.25, 0.5]))
def test_fused_qkv_scalar_scale_load_fills_all_logical_slots(self):
layer = self._make_qkv_layer()
scale = PerTensorScaleParameter(
data=torch.empty(3, dtype=torch.float32),
weight_loader=layer.weight_loader_v2,
)
layer.weight_loader_v2(scale, torch.tensor(0.125, dtype=torch.float32))
torch.testing.assert_close(scale, torch.tensor([0.125, 0.125, 0.125]))
def test_explicit_shard_scale_loads_stay_independent(self):
layer = self._make_layer()
scale = PerTensorScaleParameter(
data=torch.empty(2, dtype=torch.float32),
weight_loader=layer.weight_loader_v2,
)
layer.weight_loader_v2(scale, torch.tensor(0.25, dtype=torch.float32), 0)
layer.weight_loader_v2(scale, torch.tensor(0.5, dtype=torch.float32), 1)
torch.testing.assert_close(scale, torch.tensor([0.25, 0.5]))
if __name__ == "__main__":
unittest.main()