[diffusion] fix: fix fp8 fused tp scale loading (#28546)
This commit is contained in:
@@ -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]))
|
||||
Reference in New Issue
Block a user