diff --git a/python/sglang/multimodal_gen/runtime/layers/linear.py b/python/sglang/multimodal_gen/runtime/layers/linear.py index 7b4cce559..94affa99b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/linear.py @@ -624,7 +624,10 @@ class MergedColumnParallelLinear(ColumnParallelLinear): if loaded_shard_id is None: if isinstance(param, PerTensorScaleParameter): - if self.tp_size > 1 and loaded_weight.shape == param.data.shape: + if loaded_weight.numel() == 1 and param.data.numel() > 1: + param.data.fill_(loaded_weight.reshape(-1)[0]) + return + if loaded_weight.shape == param.data.shape: param.data.copy_(loaded_weight) return param.load_merged_column_weight(loaded_weight=loaded_weight, shard_id=0) @@ -844,6 +847,12 @@ class QKVParallelLinear(ColumnParallelLinear): ): if loaded_shard_id is None: # special case for certain models if isinstance(param, PerTensorScaleParameter): + if loaded_weight.numel() == 1 and param.data.numel() > 1: + param.data.fill_(loaded_weight.reshape(-1)[0]) + return + if loaded_weight.shape == param.data.shape: + param.data.copy_(loaded_weight) + return param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0) return elif type(param) in (RowvLLMParameter, BasevLLMParameter): diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 828b8a647..fb5e3d60a 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -493,14 +493,6 @@ else: extras=["--transformer-path", MODELOPT_FLUX1_FP8_TRANSFORMER], run_consistency_check=True, ), - _make_modelopt_ci_case( - "flux2_modelopt_fp8_t2i", - model_path=DEFAULT_FLUX_2_DEV_MODEL_NAME_FOR_TEST, - modality="image", - sampling_params=MODELOPT_T2I_CI_sampling_params, - extras=["--transformer-path", MODELOPT_FLUX2_FP8_TRANSFORMER], - run_consistency_check=True, - ), _make_modelopt_ci_case( "wan22_modelopt_fp8_t2v", model_path=MODELOPT_WAN22_FP8_MODEL, @@ -578,6 +570,18 @@ else: ONE_GPU_B200_CASES = ONE_GPU_MODELOPT_NVFP4_CASES TWO_GPU_CASES = [ + DiffusionTestCase( + "flux2_modelopt_fp8_tp2_t2i", + DiffusionServerArgs( + model_path=DEFAULT_FLUX_2_DEV_MODEL_NAME_FOR_TEST, + modality="image", + tp_size=2, + extras=["--transformer-path", MODELOPT_FLUX2_FP8_TRANSFORMER], + ), + MODELOPT_T2I_CI_sampling_params, + run_perf_check=False, + run_component_accuracy_check=False, + ), DiffusionTestCase( "ideogram4_fp8_tp2_t2i", DiffusionServerArgs( diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index ee8a79f96..49f3bce46 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -34,7 +34,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_CI_DATA_REVISION = "51a6a6cd592983e1b8dadc9d7981fac63cd02800" +SGL_TEST_FILES_CI_DATA_REVISION = "66370f48f239c08c044ac47e6eb898a01be37519" if current_platform.is_npu(): SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7" diff --git a/python/sglang/multimodal_gen/test/unit/test_parallel_linear_weight_loading.py b/python/sglang/multimodal_gen/test/unit/test_parallel_linear_weight_loading.py new file mode 100644 index 000000000..8d9deb6a2 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_parallel_linear_weight_loading.py @@ -0,0 +1,57 @@ +import pytest +import torch + +from sglang.multimodal_gen.runtime.layers.linear import ( + MergedColumnParallelLinear, + QKVParallelLinear, +) +from sglang.multimodal_gen.runtime.models.parameter import PerTensorScaleParameter + + +def _per_tensor_scale(values: list[float]) -> PerTensorScaleParameter: + return PerTensorScaleParameter( + data=torch.tensor(values, dtype=torch.float32), + weight_loader=lambda *_args, **_kwargs: None, + ) + + +@pytest.mark.parametrize("loaded_weight", [torch.tensor(0.25), torch.tensor([0.25])]) +def test_merged_column_parallel_scalar_scale_load_fills_fused_slots(loaded_weight): + layer = MergedColumnParallelLinear.__new__(MergedColumnParallelLinear) + layer.tp_size = 2 + param = _per_tensor_scale([-1.0, -2.0]) + + layer.weight_loader_v2(param, loaded_weight) + + assert torch.equal(param.data, torch.tensor([0.25, 0.25])) + + +@pytest.mark.parametrize("loaded_weight", [torch.tensor(0.5), torch.tensor([0.5])]) +def test_qkv_parallel_scalar_scale_load_fills_fused_slots(loaded_weight): + layer = QKVParallelLinear.__new__(QKVParallelLinear) + layer.tp_size = 2 + param = _per_tensor_scale([-1.0, -2.0, -3.0]) + + layer.weight_loader_v2(param, loaded_weight) + + assert torch.equal(param.data, torch.tensor([0.5, 0.5, 0.5])) + + +def test_merged_column_parallel_full_scale_vector_loads_all_fused_slots(): + layer = MergedColumnParallelLinear.__new__(MergedColumnParallelLinear) + layer.tp_size = 1 + param = _per_tensor_scale([-1.0, -2.0]) + + layer.weight_loader_v2(param, torch.tensor([0.25, 0.75])) + + assert torch.equal(param.data, torch.tensor([0.25, 0.75])) + + +def test_qkv_parallel_full_scale_vector_loads_all_fused_slots(): + layer = QKVParallelLinear.__new__(QKVParallelLinear) + layer.tp_size = 1 + param = _per_tensor_scale([-1.0, -2.0, -3.0]) + + layer.weight_loader_v2(param, torch.tensor([0.25, 0.5, 0.75])) + + assert torch.equal(param.data, torch.tensor([0.25, 0.5, 0.75]))