[diffusion] chore: optimize model weight loading (#34064)

This commit is contained in:
Mick
2026-08-10 20:07:48 +08:00
committed by GitHub
parent 4eaaeda004
commit 8ba9385097
8 changed files with 299 additions and 28 deletions
+1
View File
@@ -81,6 +81,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
- `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks.
- `--num-gpus {N}`: number of GPUs to use
- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths. `speed` keeps `torch.compile` disabled unless a model-specific deployment config opts in after validation; pass `--enable-torch-compile true` to enable it explicitly. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes.
- `--direct-gpu-weight-loading {true|false}`: opt into direct GPU loading for an unquantized, GPU-resident, TP=1 DiT by materializing its complete checkpoint state dict on GPU. Startup impact is model-dependent, so benchmark the target model before deployment. Disabled by default because checkpoint and model weights coexist temporarily, substantially increasing peak GPU memory. It is incompatible with DiT CPU/layerwise offload and FSDP.
- `--tp-size {N}`: tensor parallelism size. Depending on the pipeline, it can shard the DiT, one or more encoders, or both.
- `--sp-degree {N}`: sequence parallelism size
- `--dp-size {N}` (alias `--data-parallel-size`): number of data-parallel replicas. Each replica is a full copy of the engine on `num_gpus / N` GPUs with its own ingress; generation requests round-robin across replicas, realtime sessions stick to the replica holding their state, and control operations (weights, LoRA, memory occupation, shutdown) apply to every replica. Combines with the other parallelism axes (`num_gpus = dp × cfg × tp × sp`); monolithic serving only.
@@ -39,6 +39,17 @@ _is_npu = is_npu()
logger = init_logger(__name__)
def _resolve_checkpoint_load_device(
runtime_device: torch.device,
*,
component_cpu_offload: bool,
runtime_quant_config: object | None,
) -> torch.device:
if component_cpu_offload and runtime_quant_config is None:
return torch.device("cpu")
return runtime_device
def _default_quantized_attention_backend(
quant_spec: TransformerQuantLoadSpec, server_args: ServerArgs
) -> AttentionBackendEnum | None:
@@ -198,11 +209,30 @@ class TransformerLoader(ComponentLoader):
logger.debug("quantization config: %s", init_params["quant_config"])
local_torch_device = get_local_torch_device()
checkpoint_load_device = _resolve_checkpoint_load_device(
local_torch_device,
component_cpu_offload=bool(component_server_args.dit_cpu_offload),
runtime_quant_config=quant_spec.runtime_quant_config,
)
direct_gpu_weight_loading = bool(
component_server_args.direct_gpu_weight_loading
)
if direct_gpu_weight_loading and quant_spec.runtime_quant_config is not None:
raise ValueError(
"--direct-gpu-weight-loading supports only unquantized DiT checkpoints"
)
weight_load_plan = WeightLoadPlan.for_component(
checkpoint_load_device=local_torch_device,
checkpoint_load_device=checkpoint_load_device,
needs_device_weight_postprocess=quant_spec.needs_device_weight_postprocess,
component_cpu_offload=bool(component_server_args.dit_cpu_offload),
load_full_state_dict_on_device=direct_gpu_weight_loading,
)
if direct_gpu_weight_loading:
logger.warning(
"Direct GPU weight loading is enabled for %s; the complete checkpoint "
"state dict and materialized model weights may coexist on GPU during startup",
component_name,
)
quantized_attn_backend = _default_quantized_attention_backend(
quant_spec, component_server_args
@@ -9,6 +9,7 @@
from collections import Counter, defaultdict
from collections.abc import Callable, Generator
from itertools import chain
from types import MethodType
from typing import Any
import torch
@@ -26,7 +27,12 @@ from torch.distributed.fsdp import (
from torch.nn.modules.module import _IncompatibleKeys
from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
UnquantizedLinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
attach_bitsandbytes_4bit_quant_states,
build_bitsandbytes_4bit_quant_states,
@@ -93,6 +99,45 @@ def _make_param_like(
return new_param
def _can_assign_cpu_tensor_without_copy(
actual_param: torch.nn.Parameter,
full_tensor: torch.Tensor,
target_param: torch.Tensor,
) -> bool:
"""Return whether a TP=1 linear loader would only copy this CPU tensor."""
if full_tensor.device.type != "cpu":
return False
weight_loader = actual_param.__dict__.get("weight_loader")
if not isinstance(weight_loader, MethodType):
return False
owner = weight_loader.__self__
if not isinstance(
owner,
(ReplicatedLinear, ColumnParallelLinear, RowParallelLinear),
):
return False
if not isinstance(owner.quant_method, UnquantizedLinearMethod):
return False
if not isinstance(owner, ReplicatedLinear) and owner.tp_size != 1:
return False
if type(actual_param) is not nn.Parameter:
return False
if any(
actual_param.__dict__.get(attribute, False)
for attribute in (
"is_metadata",
"is_sharded_weight",
"needs_scalar_to_array",
)
):
return False
return (
full_tensor.shape == target_param.shape
and full_tensor.dtype == target_param.dtype
)
def _make_class_name_shard_condition(class_names: set[str]):
def shard_condition(n: str, m: nn.Module) -> bool:
return type(m).__name__ in class_names
@@ -282,7 +327,8 @@ def maybe_load_fsdp_model(
preconverted_state_dict = None
is_bnb_quantized = _is_bitsandbytes_quant_config(init_params.get("quant_config"))
if (
use_fsdp
not weight_load_plan.load_full_state_dict_on_device
and use_fsdp
and weight_dir_list
and preprocess_loaded_state_dict is None
and not is_bnb_quantized
@@ -295,7 +341,8 @@ def maybe_load_fsdp_model(
)
)
elif (
not use_fsdp
not weight_load_plan.load_full_state_dict_on_device
and not use_fsdp
and weight_dir_list
and preprocess_loaded_state_dict is None
and not is_bnb_quantized
@@ -309,7 +356,13 @@ def maybe_load_fsdp_model(
)
if preconverted_state_dict is None:
weight_iterator = safetensors_weights_iterator(weight_dir_list)
if weight_load_plan.load_full_state_dict_on_device:
weight_iterator = safetensors_weights_iterator(
weight_dir_list,
weight_load_plan=weight_load_plan,
)
else:
weight_iterator = safetensors_weights_iterator(weight_dir_list)
if preprocess_loaded_state_dict is not None:
weight_iterator = preprocess_loaded_state_dict(weight_iterator)
if is_bnb_quantized:
@@ -611,7 +664,8 @@ def load_model_from_full_model_state_dict(
sharded_tensor = sharded_tensor.cpu()
elif not isinstance(meta_sharded_param, dist_tensor.DTensor):
full_tensor = full_tensor.to(
device=checkpoint_load_device, dtype=target_dtype
device=checkpoint_load_device,
dtype=target_dtype,
)
actual_param = rank_local_checkpoint.get_param_for_weight_loading(
model, param_dict, target_param_name
@@ -623,30 +677,38 @@ def load_model_from_full_model_state_dict(
)
if weight_loader is not None:
assert actual_param is not None
sharded_tensor = torch.empty_like(
if _can_assign_cpu_tensor_without_copy(
actual_param,
full_tensor,
meta_sharded_param,
device=checkpoint_load_device,
dtype=target_dtype,
)
# Preserve requires_grad flag to avoid errors with non-floating dtypes
requires_grad = getattr(meta_sharded_param, "requires_grad", False)
temp_param = _make_param_like(actual_param, sharded_tensor)
if not (
sharded_tensor.is_floating_point() or sharded_tensor.is_complex()
):
requires_grad = False
temp_param.requires_grad = requires_grad
try:
weight_loader(temp_param, full_tensor)
except AssertionError as exc:
raise AssertionError(
"Failed to shard/load parameter "
f"{target_param_name}: full_tensor.shape={tuple(full_tensor.shape)}, "
f"meta_sharded_param.shape={tuple(meta_sharded_param.shape)}, "
f"temp_param.shape={tuple(temp_param.shape)}, "
f"param_cls={type(actual_param).__name__}"
) from exc
sharded_tensor = temp_param.data
sharded_tensor = full_tensor
else:
sharded_tensor = torch.empty_like(
meta_sharded_param,
device=checkpoint_load_device,
dtype=target_dtype,
)
# Preserve requires_grad flag to avoid errors with non-floating dtypes
requires_grad = meta_sharded_param.requires_grad
temp_param = _make_param_like(actual_param, sharded_tensor)
if not (
sharded_tensor.is_floating_point()
or sharded_tensor.is_complex()
):
requires_grad = False
temp_param.requires_grad = requires_grad
try:
weight_loader(temp_param, full_tensor)
except AssertionError as exc:
raise AssertionError(
"Failed to shard/load parameter "
f"{target_param_name}: full_tensor.shape={tuple(full_tensor.shape)}, "
f"meta_sharded_param.shape={tuple(meta_sharded_param.shape)}, "
f"temp_param.shape={tuple(temp_param.shape)}, "
f"param_cls={type(actual_param).__name__}"
) from exc
sharded_tensor = temp_param.data
else:
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
sharded_tensor = full_tensor
@@ -13,6 +13,8 @@ class WeightLoadPlan:
weight_postprocess_device: torch.device | None = None
# Delay non-FSDP component CPU offload until after weight postprocessing.
defer_component_cpu_offload: bool = False
# keep the complete mapped checkpoint state dict on the load device
load_full_state_dict_on_device: bool = False
@classmethod
def for_component(
@@ -21,6 +23,7 @@ class WeightLoadPlan:
checkpoint_load_device: torch.device,
needs_device_weight_postprocess: bool,
component_cpu_offload: bool,
load_full_state_dict_on_device: bool = False,
) -> "WeightLoadPlan":
# if on-device weight postprocessing is required, load directly to device to speedup loading
weight_postprocess_device = (
@@ -32,4 +35,5 @@ class WeightLoadPlan:
defer_component_cpu_offload=(
needs_device_weight_postprocess and component_cpu_offload
),
load_full_state_dict_on_device=load_full_state_dict_on_device,
)
@@ -299,6 +299,8 @@ class ServerArgs(DisaggServerArgsMixin):
# CPU offload parameters
dit_cpu_offload: bool | None = None
# trade checkpoint-loading peak memory for faster ordinary DiT startup
direct_gpu_weight_loading: bool = False
# if true, select the DiT layerwise group
dit_layerwise_offload: bool | None = None
layerwise_offload_components: list[str] | None = None
@@ -508,6 +510,7 @@ class ServerArgs(DisaggServerArgsMixin):
self._validate_scheduler_rpc_timeout()
self._validate_pipeline()
self._validate_offload()
self._validate_direct_gpu_weight_loading()
if not current_platform.is_cpu():
self._validate_parallelism()
self._validate_cfg_parallel()
@@ -1771,6 +1774,15 @@ class ServerArgs(DisaggServerArgsMixin):
action=StoreBoolean,
help="Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
)
parser.add_argument(
"--direct-gpu-weight-loading",
action=StoreBoolean,
default=ServerArgs.direct_gpu_weight_loading,
help="Load the full unquantized DiT checkpoint state dict directly "
"onto GPU before assigning model parameters. This may reduce startup "
"time depending on the model, but temporarily requires checkpoint "
"weights and model weights to coexist on GPU. Disabled by default.",
)
parser.add_argument(
"--dit-layerwise-offload",
action=StoreBoolean,
@@ -2634,6 +2646,23 @@ class ServerArgs(DisaggServerArgsMixin):
"--performance-mode speed for GPU-resident defaults when memory allows."
)
def _validate_direct_gpu_weight_loading(self) -> None:
if not self.direct_gpu_weight_loading:
return
if not current_platform.is_cuda():
raise ValueError("--direct-gpu-weight-loading requires CUDA")
if self.dit_cpu_offload or self.is_dit_layerwise_offload_selected:
raise ValueError(
"--direct-gpu-weight-loading requires a GPU-resident DiT; disable "
"DiT CPU and layerwise offload"
)
if self.use_fsdp_inference:
raise ValueError(
"--direct-gpu-weight-loading does not support FSDP inference"
)
if self.tp_size != 1:
raise ValueError("--direct-gpu-weight-loading requires --tp-size 1")
def _validate_parallelism(self):
if self.kv_gather_degree < 1:
raise ValueError("kv_gather_degree must be >= 1")
@@ -9,7 +9,9 @@ import torch
from safetensors.torch import safe_open, save_file
from torch import nn
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
from sglang.multimodal_gen.runtime.loader import fsdp_load, rank_local_checkpoint
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
class _UniformDtypeModel(nn.Module):
@@ -26,6 +28,12 @@ class _MixedDtypeModel(_UniformDtypeModel):
_fsdp_mixed_dtype_params = True
class _ReplicatedLinearModel(_UniformDtypeModel):
def __init__(self) -> None:
super().__init__()
self.proj = ReplicatedLinear(4, 4, bias=False)
class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
def _load_and_capture_policy(
self,
@@ -97,6 +105,61 @@ class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
shard_model.assert_not_called()
class TestOrdinaryWeightLoading(unittest.TestCase):
def test_direct_device_loading_skips_rank_local_cpu_checkpoint(self):
load_plan = WeightLoadPlan(
checkpoint_load_device=torch.device("cuda:0"),
load_full_state_dict_on_device=True,
)
with (
patch.object(fsdp_load.current_platform, "is_mps", return_value=False),
patch.object(
rank_local_checkpoint,
"try_load_rank_local_tp_state_dict",
) as rank_local_load,
patch.object(
fsdp_load,
"safetensors_weights_iterator",
return_value=iter(()),
) as weight_iterator,
patch.object(fsdp_load, "load_model_from_full_model_state_dict"),
):
fsdp_load.maybe_load_fsdp_model(
model_cls=_UniformDtypeModel,
init_params={},
weight_dir_list=["model.safetensors"],
device=torch.device("cuda:0"),
hsdp_replicate_dim=1,
hsdp_shard_dim=1,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
weight_load_plan=load_plan,
)
rank_local_load.assert_not_called()
weight_iterator.assert_called_once_with(
["model.safetensors"],
weight_load_plan=load_plan,
)
def test_tp1_unquantized_linear_assigns_checkpoint_tensor_without_copy(self):
with torch.device("meta"):
model = _ReplicatedLinearModel()
checkpoint_weight = torch.arange(16, dtype=torch.float32).reshape(4, 4)
fsdp_load.load_model_from_full_model_state_dict(
model,
iter((("proj.weight", checkpoint_weight),)),
checkpoint_load_device=torch.device("cpu"),
param_dtype=torch.float32,
strict=True,
param_names_mapping=fsdp_load.get_param_names_mapping({}),
)
self.assertEqual(model.proj.weight.data_ptr(), checkpoint_weight.data_ptr())
torch.testing.assert_close(model.proj.weight, checkpoint_weight)
class TestRankLocalSafetensorsRead(unittest.TestCase):
def _source(
self,
@@ -57,6 +57,7 @@ from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import (
MiniMaxH3Pipeline,
)
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import (
MAX_SCHEDULER_RPC_TIMEOUT_S,
ServerArgs,
@@ -2302,5 +2303,44 @@ class TestNcclNvlsArgs(unittest.TestCase):
self.assertFalse(disabled_args.enable_nccl_nvls)
class TestDirectGpuWeightLoading(unittest.TestCase):
def _args(self) -> ServerArgs:
args = ServerArgs.__new__(ServerArgs)
args.direct_gpu_weight_loading = True
args.dit_cpu_offload = False
args.layerwise_offload_components = []
args.use_fsdp_inference = False
args.tp_size = 1
return args
def test_cli_defaults_off_and_parses_explicit_enable(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
default_args, _ = parser.parse_known_args(["--model-path", "/fake"])
enabled_args, _ = parser.parse_known_args(
["--model-path", "/fake", "--direct-gpu-weight-loading"]
)
self.assertFalse(default_args.direct_gpu_weight_loading)
self.assertTrue(enabled_args.direct_gpu_weight_loading)
def test_rejects_cpu_offload_fsdp_and_tp(self):
cpu_offload_args = self._args()
cpu_offload_args.dit_cpu_offload = True
fsdp_args = self._args()
fsdp_args.use_fsdp_inference = True
tp_args = self._args()
tp_args.tp_size = 2
with patch.object(current_platform, "is_cuda", return_value=True):
with self.assertRaisesRegex(ValueError, "GPU-resident DiT"):
cpu_offload_args._validate_direct_gpu_weight_loading()
with self.assertRaisesRegex(ValueError, "FSDP"):
fsdp_args._validate_direct_gpu_weight_loading()
with self.assertRaisesRegex(ValueError, "tp-size 1"):
tp_args._validate_direct_gpu_weight_loading()
if __name__ == "__main__":
unittest.main()
@@ -58,6 +58,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
from sglang.multimodal_gen.runtime.loader.component_loaders import transformer_loader
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
_default_quantized_attention_backend,
_resolve_checkpoint_load_device,
_warn_if_expected_param_dtype_missing,
)
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
@@ -118,6 +119,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
quantization_ignored_layers=None,
tp_size=1,
dit_cpu_offload=False,
direct_gpu_weight_loading=False,
text_encoder_cpu_offload=False,
)
defaults.update(overrides)
@@ -253,6 +255,46 @@ class TestTransformerQuantHelpers(unittest.TestCase):
self.assertEqual(plan.checkpoint_load_device, device)
self.assertEqual(plan.weight_postprocess_device, device)
self.assertTrue(plan.defer_component_cpu_offload)
self.assertFalse(plan.load_full_state_dict_on_device)
def test_weight_load_plan_can_keep_full_state_dict_on_device(self):
plan = WeightLoadPlan.for_component(
checkpoint_load_device=torch.device("cuda:0"),
needs_device_weight_postprocess=False,
component_cpu_offload=False,
load_full_state_dict_on_device=True,
)
self.assertTrue(plan.load_full_state_dict_on_device)
def test_unquantized_cpu_offload_loads_checkpoint_on_cpu(self):
device = _resolve_checkpoint_load_device(
torch.device("cuda:0"),
component_cpu_offload=True,
runtime_quant_config=None,
)
self.assertEqual(device, torch.device("cpu"))
def test_quantized_cpu_offload_keeps_checkpoint_on_runtime_device(self):
runtime_device = torch.device("cuda:0")
device = _resolve_checkpoint_load_device(
runtime_device,
component_cpu_offload=True,
runtime_quant_config=object(),
)
self.assertEqual(device, runtime_device)
def test_resident_transformer_loads_checkpoint_on_runtime_device(self):
runtime_device = torch.device("cuda:0")
device = _resolve_checkpoint_load_device(
runtime_device,
component_cpu_offload=False,
runtime_quant_config=None,
)
self.assertEqual(device, runtime_device)
def test_mixed_model_with_expected_dtype_does_not_warn(self):
model = torch.nn.Module()