From b24bd44556a4b8bc0134bf1d5ac64b42181f7f0f Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 29 Aug 2026 22:43:25 +0800 Subject: [PATCH] [diffusion] feat: avoid direct GPU parameter copies (#36832) --- docs/docs/sglang-diffusion/api/cli.mdx | 2 +- .../component_loaders/transformer_loader.py | 5 +- .../runtime/loader/fsdp_load.py | 24 ++++-- .../test/unit/test_fsdp_load.py | 86 +++++++++++++++---- 4 files changed, 91 insertions(+), 26 deletions(-) diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index 1eb742c02..a62b95079 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -86,7 +86,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 dispatches residency from selected-GPU headroom and workload type: image DiTs stay resident above the 45 GiB threshold, while video DiT placement remains model-specific. It uses 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. +- `--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. Compatible tensors become parameter storage without another GPU copy, while tensors requiring transforms can still need temporary allocations. It is disabled by default; startup time and peak GPU memory remain model-dependent, so benchmark the target model before deployment. 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. diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 723cf2899..93fabd89d 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -427,8 +427,9 @@ class TransformerLoader(ComponentLoader): ) 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", + "Direct GPU weight loading is enabled for %s; compatible checkpoint " + "tensors become model storage, while transformed tensors may still " + "require temporary GPU allocations", component_name, ) diff --git a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py index 013dfae83..ca8ec495b 100644 --- a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py +++ b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py @@ -98,14 +98,12 @@ def _make_param_like( return new_param -def _can_assign_cpu_tensor_without_copy( +def _can_assign_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 + """Return whether a TP=1 linear loader would only copy this tensor.""" weight_loader = actual_param.__dict__.get("weight_loader") if not isinstance(weight_loader, MethodType): return False @@ -134,6 +132,8 @@ def _can_assign_cpu_tensor_without_copy( return ( full_tensor.shape == target_param.shape and full_tensor.dtype == target_param.dtype + and full_tensor.layout == target_param.layout + and full_tensor.stride() == target_param.stride() ) @@ -432,6 +432,9 @@ def maybe_load_fsdp_model( cpu_offload=load_on_cpu, param_names_mapping=param_names_mapping_fn, keep_checkpoint_mapping=keep_checkpoint_mapping, + allow_device_tensor_assignment=( + weight_load_plan.load_full_state_dict_on_device + ), preconverted_state_dict=preconverted_state_dict, ) if bnb_quant_states: @@ -568,6 +571,7 @@ def load_model_from_full_model_state_dict( ] | None ) = None, + allow_device_tensor_assignment: bool = False, ) -> _IncompatibleKeys: """ Converting full state dict into a sharded state dict @@ -581,6 +585,10 @@ def load_model_from_full_model_state_dict( cpu_offload (bool): flag to check if FSDP offload is enabled param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name keep_checkpoint_mapping (bool): retain compatible CPU checkpoint tensors instead of copying them + allow_device_tensor_assignment (bool): adopt compatible checkpoint tensors + already materialized on the target device. This is reserved for an + explicit full-state direct-device load; ordinary loading keeps its + established parameter materialization path. Returns: ``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields: * **missing_keys** is a list of str containing the missing keys @@ -734,10 +742,10 @@ def load_model_from_full_model_state_dict( sharded_tensor = full_tensor elif weight_loader is not None: assert actual_param is not None - if _can_assign_cpu_tensor_without_copy( - actual_param, - full_tensor, - meta_sharded_param, + if ( + full_tensor.device.type == "cpu" or allow_device_tensor_assignment + ) and _can_assign_tensor_without_copy( + actual_param, full_tensor, meta_sharded_param ): sharded_tensor = full_tensor else: diff --git a/python/sglang/multimodal_gen/test/unit/test_fsdp_load.py b/python/sglang/multimodal_gen/test/unit/test_fsdp_load.py index b2022cf11..5cd84a4b5 100644 --- a/python/sglang/multimodal_gen/test/unit/test_fsdp_load.py +++ b/python/sglang/multimodal_gen/test/unit/test_fsdp_load.py @@ -9,7 +9,11 @@ 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.layers.linear import ( + ColumnParallelLinear, + ReplicatedLinear, + UnquantizedLinearMethod, +) from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import ( BitsAndBytesConfig, ) @@ -180,6 +184,26 @@ class TestFSDPEntrypointRegistration(unittest.TestCase): class TestOrdinaryWeightLoading(unittest.TestCase): + def _load_replicated_weight( + self, device: torch.device, *, allow_device_tensor_assignment: bool = False + ) -> tuple[torch.nn.Parameter, torch.Tensor]: + with torch.device("meta"): + model = _ReplicatedLinearModel() + checkpoint_weight = torch.arange( + 16, dtype=torch.float32, device=device + ).reshape(4, 4) + + fsdp_load.load_model_from_full_model_state_dict( + model, + iter((("proj.weight", checkpoint_weight),)), + checkpoint_load_device=device, + param_dtype=torch.float32, + strict=True, + param_names_mapping=fsdp_load.get_param_names_mapping({}), + allow_device_tensor_assignment=allow_device_tensor_assignment, + ) + return model.proj.weight, checkpoint_weight + def test_direct_device_loading_skips_rank_local_cpu_checkpoint(self): load_plan = WeightLoadPlan( checkpoint_load_device=torch.device("cuda:0"), @@ -217,22 +241,54 @@ class TestOrdinaryWeightLoading(unittest.TestCase): 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({}), + def test_tp1_unquantized_linear_adopts_cpu_checkpoint_storage(self): + model_weight, checkpoint_weight = self._load_replicated_weight( + torch.device("cpu") ) - self.assertEqual(model.proj.weight.data_ptr(), checkpoint_weight.data_ptr()) - torch.testing.assert_close(model.proj.weight, checkpoint_weight) + self.assertEqual(model_weight.data_ptr(), checkpoint_weight.data_ptr()) + torch.testing.assert_close(model_weight, checkpoint_weight) + + @unittest.skipUnless(torch.cuda.is_available(), "CUDA is required") + def test_ordinary_cuda_loading_preserves_materialization_path(self): + model_weight, checkpoint_weight = self._load_replicated_weight( + torch.device("cuda:0") + ) + + self.assertNotEqual(model_weight.data_ptr(), checkpoint_weight.data_ptr()) + torch.testing.assert_close(model_weight, checkpoint_weight) + + @unittest.skipUnless(torch.cuda.is_available(), "CUDA is required") + def test_direct_cuda_loading_adopts_checkpoint_storage(self): + model_weight, checkpoint_weight = self._load_replicated_weight( + torch.device("cuda:0"), allow_device_tensor_assignment=True + ) + + self.assertEqual(model_weight.data_ptr(), checkpoint_weight.data_ptr()) + torch.testing.assert_close(model_weight, checkpoint_weight) + + def test_zero_copy_assignment_rejects_incompatible_layout_or_tp_weights(self): + with torch.device("meta"): + model = _ReplicatedLinearModel() + param = model.proj.weight + tensor = torch.empty(4, 4) + + self.assertFalse( + fsdp_load._can_assign_tensor_without_copy( + param, tensor.as_strided((4, 4), (1, 4)), param + ) + ) + + tp_owner = ColumnParallelLinear.__new__(ColumnParallelLinear) + nn.Module.__init__(tp_owner) + tp_owner.quant_method = UnquantizedLinearMethod() + tp_owner.tp_size = 1 + tp_param = nn.Parameter(tensor) + tp_param.weight_loader = tp_owner.weight_loader + tp_owner.tp_size = 2 + self.assertFalse( + fsdp_load._can_assign_tensor_without_copy(tp_param, tensor, tp_param) + ) class TestDevicePostprocessMove(unittest.TestCase):