[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 loaded_shard_id is None:
|
||||||
if isinstance(param, PerTensorScaleParameter):
|
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)
|
param.data.copy_(loaded_weight)
|
||||||
return
|
return
|
||||||
param.load_merged_column_weight(loaded_weight=loaded_weight, shard_id=0)
|
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 loaded_shard_id is None: # special case for certain models
|
||||||
if isinstance(param, PerTensorScaleParameter):
|
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)
|
param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0)
|
||||||
return
|
return
|
||||||
elif type(param) in (RowvLLMParameter, BasevLLMParameter):
|
elif type(param) in (RowvLLMParameter, BasevLLMParameter):
|
||||||
|
|||||||
@@ -493,14 +493,6 @@ else:
|
|||||||
extras=["--transformer-path", MODELOPT_FLUX1_FP8_TRANSFORMER],
|
extras=["--transformer-path", MODELOPT_FLUX1_FP8_TRANSFORMER],
|
||||||
run_consistency_check=True,
|
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(
|
_make_modelopt_ci_case(
|
||||||
"wan22_modelopt_fp8_t2v",
|
"wan22_modelopt_fp8_t2v",
|
||||||
model_path=MODELOPT_WAN22_FP8_MODEL,
|
model_path=MODELOPT_WAN22_FP8_MODEL,
|
||||||
@@ -578,6 +570,18 @@ else:
|
|||||||
ONE_GPU_B200_CASES = ONE_GPU_MODELOPT_NVFP4_CASES
|
ONE_GPU_B200_CASES = ONE_GPU_MODELOPT_NVFP4_CASES
|
||||||
|
|
||||||
TWO_GPU_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(
|
DiffusionTestCase(
|
||||||
"ideogram4_fp8_tp2_t2i",
|
"ideogram4_fp8_tp2_t2i",
|
||||||
DiffusionServerArgs(
|
DiffusionServerArgs(
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "51a6a6cd592983e1b8dadc9d7981fac63cd02800"
|
SGL_TEST_FILES_CI_DATA_REVISION = "66370f48f239c08c044ac47e6eb898a01be37519"
|
||||||
|
|
||||||
if current_platform.is_npu():
|
if current_platform.is_npu():
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7"
|
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