[diffusion] fix: fix fp8 fused tp scale loading (#28546)

This commit is contained in:
Mick
2026-06-18 14:38:38 +08:00
committed by GitHub
parent 5d1949152d
commit 3b61dc32c9
4 changed files with 80 additions and 10 deletions
@@ -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):
@@ -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(
@@ -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"
@@ -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]))