[diffusion] chore: optimize model weight loading (#34064)
This commit is contained in:
@@ -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.
|
- `--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
|
- `--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.
|
- `--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.
|
- `--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
|
- `--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.
|
- `--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.
|
||||||
|
|||||||
+31
-1
@@ -39,6 +39,17 @@ _is_npu = is_npu()
|
|||||||
logger = init_logger(__name__)
|
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(
|
def _default_quantized_attention_backend(
|
||||||
quant_spec: TransformerQuantLoadSpec, server_args: ServerArgs
|
quant_spec: TransformerQuantLoadSpec, server_args: ServerArgs
|
||||||
) -> AttentionBackendEnum | None:
|
) -> AttentionBackendEnum | None:
|
||||||
@@ -198,10 +209,29 @@ class TransformerLoader(ComponentLoader):
|
|||||||
logger.debug("quantization config: %s", init_params["quant_config"])
|
logger.debug("quantization config: %s", init_params["quant_config"])
|
||||||
|
|
||||||
local_torch_device = get_local_torch_device()
|
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(
|
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,
|
needs_device_weight_postprocess=quant_spec.needs_device_weight_postprocess,
|
||||||
component_cpu_offload=bool(component_server_args.dit_cpu_offload),
|
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(
|
quantized_attn_backend = _default_quantized_attention_backend(
|
||||||
|
|||||||
@@ -9,6 +9,7 @@
|
|||||||
from collections import Counter, defaultdict
|
from collections import Counter, defaultdict
|
||||||
from collections.abc import Callable, Generator
|
from collections.abc import Callable, Generator
|
||||||
from itertools import chain
|
from itertools import chain
|
||||||
|
from types import MethodType
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -26,7 +27,12 @@ from torch.distributed.fsdp import (
|
|||||||
from torch.nn.modules.module import _IncompatibleKeys
|
from torch.nn.modules.module import _IncompatibleKeys
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in
|
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 (
|
from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
|
||||||
attach_bitsandbytes_4bit_quant_states,
|
attach_bitsandbytes_4bit_quant_states,
|
||||||
build_bitsandbytes_4bit_quant_states,
|
build_bitsandbytes_4bit_quant_states,
|
||||||
@@ -93,6 +99,45 @@ def _make_param_like(
|
|||||||
return new_param
|
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 _make_class_name_shard_condition(class_names: set[str]):
|
||||||
def shard_condition(n: str, m: nn.Module) -> bool:
|
def shard_condition(n: str, m: nn.Module) -> bool:
|
||||||
return type(m).__name__ in class_names
|
return type(m).__name__ in class_names
|
||||||
@@ -282,7 +327,8 @@ def maybe_load_fsdp_model(
|
|||||||
preconverted_state_dict = None
|
preconverted_state_dict = None
|
||||||
is_bnb_quantized = _is_bitsandbytes_quant_config(init_params.get("quant_config"))
|
is_bnb_quantized = _is_bitsandbytes_quant_config(init_params.get("quant_config"))
|
||||||
if (
|
if (
|
||||||
use_fsdp
|
not weight_load_plan.load_full_state_dict_on_device
|
||||||
|
and use_fsdp
|
||||||
and weight_dir_list
|
and weight_dir_list
|
||||||
and preprocess_loaded_state_dict is None
|
and preprocess_loaded_state_dict is None
|
||||||
and not is_bnb_quantized
|
and not is_bnb_quantized
|
||||||
@@ -295,7 +341,8 @@ def maybe_load_fsdp_model(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
elif (
|
elif (
|
||||||
not use_fsdp
|
not weight_load_plan.load_full_state_dict_on_device
|
||||||
|
and not use_fsdp
|
||||||
and weight_dir_list
|
and weight_dir_list
|
||||||
and preprocess_loaded_state_dict is None
|
and preprocess_loaded_state_dict is None
|
||||||
and not is_bnb_quantized
|
and not is_bnb_quantized
|
||||||
@@ -309,6 +356,12 @@ def maybe_load_fsdp_model(
|
|||||||
)
|
)
|
||||||
|
|
||||||
if preconverted_state_dict is None:
|
if preconverted_state_dict is None:
|
||||||
|
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)
|
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||||
if preprocess_loaded_state_dict is not None:
|
if preprocess_loaded_state_dict is not None:
|
||||||
weight_iterator = preprocess_loaded_state_dict(weight_iterator)
|
weight_iterator = preprocess_loaded_state_dict(weight_iterator)
|
||||||
@@ -611,7 +664,8 @@ def load_model_from_full_model_state_dict(
|
|||||||
sharded_tensor = sharded_tensor.cpu()
|
sharded_tensor = sharded_tensor.cpu()
|
||||||
elif not isinstance(meta_sharded_param, dist_tensor.DTensor):
|
elif not isinstance(meta_sharded_param, dist_tensor.DTensor):
|
||||||
full_tensor = full_tensor.to(
|
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(
|
actual_param = rank_local_checkpoint.get_param_for_weight_loading(
|
||||||
model, param_dict, target_param_name
|
model, param_dict, target_param_name
|
||||||
@@ -623,16 +677,24 @@ def load_model_from_full_model_state_dict(
|
|||||||
)
|
)
|
||||||
if weight_loader is not None:
|
if weight_loader is not None:
|
||||||
assert actual_param is not None
|
assert actual_param is not None
|
||||||
|
if _can_assign_cpu_tensor_without_copy(
|
||||||
|
actual_param,
|
||||||
|
full_tensor,
|
||||||
|
meta_sharded_param,
|
||||||
|
):
|
||||||
|
sharded_tensor = full_tensor
|
||||||
|
else:
|
||||||
sharded_tensor = torch.empty_like(
|
sharded_tensor = torch.empty_like(
|
||||||
meta_sharded_param,
|
meta_sharded_param,
|
||||||
device=checkpoint_load_device,
|
device=checkpoint_load_device,
|
||||||
dtype=target_dtype,
|
dtype=target_dtype,
|
||||||
)
|
)
|
||||||
# Preserve requires_grad flag to avoid errors with non-floating dtypes
|
# Preserve requires_grad flag to avoid errors with non-floating dtypes
|
||||||
requires_grad = getattr(meta_sharded_param, "requires_grad", False)
|
requires_grad = meta_sharded_param.requires_grad
|
||||||
temp_param = _make_param_like(actual_param, sharded_tensor)
|
temp_param = _make_param_like(actual_param, sharded_tensor)
|
||||||
if not (
|
if not (
|
||||||
sharded_tensor.is_floating_point() or sharded_tensor.is_complex()
|
sharded_tensor.is_floating_point()
|
||||||
|
or sharded_tensor.is_complex()
|
||||||
):
|
):
|
||||||
requires_grad = False
|
requires_grad = False
|
||||||
temp_param.requires_grad = requires_grad
|
temp_param.requires_grad = requires_grad
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ class WeightLoadPlan:
|
|||||||
weight_postprocess_device: torch.device | None = None
|
weight_postprocess_device: torch.device | None = None
|
||||||
# Delay non-FSDP component CPU offload until after weight postprocessing.
|
# Delay non-FSDP component CPU offload until after weight postprocessing.
|
||||||
defer_component_cpu_offload: bool = False
|
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
|
@classmethod
|
||||||
def for_component(
|
def for_component(
|
||||||
@@ -21,6 +23,7 @@ class WeightLoadPlan:
|
|||||||
checkpoint_load_device: torch.device,
|
checkpoint_load_device: torch.device,
|
||||||
needs_device_weight_postprocess: bool,
|
needs_device_weight_postprocess: bool,
|
||||||
component_cpu_offload: bool,
|
component_cpu_offload: bool,
|
||||||
|
load_full_state_dict_on_device: bool = False,
|
||||||
) -> "WeightLoadPlan":
|
) -> "WeightLoadPlan":
|
||||||
# if on-device weight postprocessing is required, load directly to device to speedup loading
|
# if on-device weight postprocessing is required, load directly to device to speedup loading
|
||||||
weight_postprocess_device = (
|
weight_postprocess_device = (
|
||||||
@@ -32,4 +35,5 @@ class WeightLoadPlan:
|
|||||||
defer_component_cpu_offload=(
|
defer_component_cpu_offload=(
|
||||||
needs_device_weight_postprocess and 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
|
# CPU offload parameters
|
||||||
dit_cpu_offload: bool | None = None
|
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
|
# if true, select the DiT layerwise group
|
||||||
dit_layerwise_offload: bool | None = None
|
dit_layerwise_offload: bool | None = None
|
||||||
layerwise_offload_components: list[str] | None = None
|
layerwise_offload_components: list[str] | None = None
|
||||||
@@ -508,6 +510,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
self._validate_scheduler_rpc_timeout()
|
self._validate_scheduler_rpc_timeout()
|
||||||
self._validate_pipeline()
|
self._validate_pipeline()
|
||||||
self._validate_offload()
|
self._validate_offload()
|
||||||
|
self._validate_direct_gpu_weight_loading()
|
||||||
if not current_platform.is_cpu():
|
if not current_platform.is_cpu():
|
||||||
self._validate_parallelism()
|
self._validate_parallelism()
|
||||||
self._validate_cfg_parallel()
|
self._validate_cfg_parallel()
|
||||||
@@ -1771,6 +1774,15 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
action=StoreBoolean,
|
action=StoreBoolean,
|
||||||
help="Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
|
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(
|
parser.add_argument(
|
||||||
"--dit-layerwise-offload",
|
"--dit-layerwise-offload",
|
||||||
action=StoreBoolean,
|
action=StoreBoolean,
|
||||||
@@ -2634,6 +2646,23 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
"--performance-mode speed for GPU-resident defaults when memory allows."
|
"--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):
|
def _validate_parallelism(self):
|
||||||
if self.kv_gather_degree < 1:
|
if self.kv_gather_degree < 1:
|
||||||
raise ValueError("kv_gather_degree must be >= 1")
|
raise ValueError("kv_gather_degree must be >= 1")
|
||||||
|
|||||||
@@ -9,7 +9,9 @@ import torch
|
|||||||
from safetensors.torch import safe_open, save_file
|
from safetensors.torch import safe_open, save_file
|
||||||
from torch import nn
|
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 import fsdp_load, rank_local_checkpoint
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
|
||||||
|
|
||||||
|
|
||||||
class _UniformDtypeModel(nn.Module):
|
class _UniformDtypeModel(nn.Module):
|
||||||
@@ -26,6 +28,12 @@ class _MixedDtypeModel(_UniformDtypeModel):
|
|||||||
_fsdp_mixed_dtype_params = True
|
_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):
|
class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
|
||||||
def _load_and_capture_policy(
|
def _load_and_capture_policy(
|
||||||
self,
|
self,
|
||||||
@@ -97,6 +105,61 @@ class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
|
|||||||
shard_model.assert_not_called()
|
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):
|
class TestRankLocalSafetensorsRead(unittest.TestCase):
|
||||||
def _source(
|
def _source(
|
||||||
self,
|
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 (
|
from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import (
|
||||||
MiniMaxH3Pipeline,
|
MiniMaxH3Pipeline,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import (
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
MAX_SCHEDULER_RPC_TIMEOUT_S,
|
MAX_SCHEDULER_RPC_TIMEOUT_S,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -2302,5 +2303,44 @@ class TestNcclNvlsArgs(unittest.TestCase):
|
|||||||
self.assertFalse(disabled_args.enable_nccl_nvls)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.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 import transformer_loader
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
|
||||||
_default_quantized_attention_backend,
|
_default_quantized_attention_backend,
|
||||||
|
_resolve_checkpoint_load_device,
|
||||||
_warn_if_expected_param_dtype_missing,
|
_warn_if_expected_param_dtype_missing,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
||||||
@@ -118,6 +119,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
quantization_ignored_layers=None,
|
quantization_ignored_layers=None,
|
||||||
tp_size=1,
|
tp_size=1,
|
||||||
dit_cpu_offload=False,
|
dit_cpu_offload=False,
|
||||||
|
direct_gpu_weight_loading=False,
|
||||||
text_encoder_cpu_offload=False,
|
text_encoder_cpu_offload=False,
|
||||||
)
|
)
|
||||||
defaults.update(overrides)
|
defaults.update(overrides)
|
||||||
@@ -253,6 +255,46 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
self.assertEqual(plan.checkpoint_load_device, device)
|
self.assertEqual(plan.checkpoint_load_device, device)
|
||||||
self.assertEqual(plan.weight_postprocess_device, device)
|
self.assertEqual(plan.weight_postprocess_device, device)
|
||||||
self.assertTrue(plan.defer_component_cpu_offload)
|
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):
|
def test_mixed_model_with_expected_dtype_does_not_warn(self):
|
||||||
model = torch.nn.Module()
|
model = torch.nn.Module()
|
||||||
|
|||||||
Reference in New Issue
Block a user