[diffusion] refactor: consolidate diffusion weight load planning (#30118)

This commit is contained in:
Mick
2026-07-05 11:59:50 +08:00
committed by GitHub
parent 5e6f49c986
commit a37bc2456d
7 changed files with 209 additions and 62 deletions
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
)
from sglang.multimodal_gen.runtime.loader.fsdp_load import maybe_load_fsdp_model
from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
@@ -78,12 +79,13 @@ class BridgeLoader(ComponentLoader):
if server_args.use_fsdp_inference or (
server_args.hsdp_shard_dim is not None and fsdp_shard_conditions
):
local_torch_device = get_local_torch_device()
# Load with FSDP support
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={"config": bridge_config, "hf_config": hf_config},
weight_dir_list=safetensors_list,
device=get_local_torch_device(),
device=local_torch_device,
hsdp_replicate_dim=server_args.hsdp_replicate_dim,
hsdp_shard_dim=server_args.hsdp_shard_dim,
cpu_offload=server_args.dit_cpu_offload,
@@ -93,6 +95,9 @@ class BridgeLoader(ComponentLoader):
reduce_dtype=torch.float32,
output_dtype=None,
strict=False,
weight_load_plan=WeightLoadPlan(
checkpoint_load_device=local_torch_device
),
)
else:
# Fallback to simple loading (for non-FSDP or legacy models)
@@ -173,6 +173,49 @@ class ComponentLoader(ABC):
return component
return component.to(get_local_torch_device())
def _load_customized_with_context(
self,
component_model_path: str,
server_args: ServerArgs,
component_name: str,
attn_backend: Any,
component_attn_name: str | None,
) -> AutoModel:
with component_attn_backend_context_manager(
attn_backend, component_name=component_attn_name
):
load_kwargs = self.customized_load_kwargs_for_component(
server_args, component_name
)
component = self.load_customized(
component_model_path, server_args, component_name, **load_kwargs
)
return self._maybe_configure_layerwise_after_startup_cpu_staging(
component, server_args, component_name, load_kwargs
)
def _load_native_with_context(
self,
component_model_path: str,
server_args: ServerArgs,
component_name: str,
transformers_or_diffusers: str,
attn_backend: Any,
component_attn_name: str | None,
) -> AutoModel:
with component_attn_backend_context_manager(
attn_backend, component_name=component_attn_name
):
component = self.load_native(
component_model_path,
server_args,
transformers_or_diffusers,
component_name,
)
should_offload = self.should_offload(server_args)
target_device = self.target_device(should_offload)
return component.to(device=target_device)
def load(
self,
component_model_path: str,
@@ -209,19 +252,13 @@ class ComponentLoader(ABC):
matched_backend_key,
)
try:
with component_attn_backend_context_manager(
attn_backend, component_name=component_attn_name
):
load_kwargs = self.customized_load_kwargs_for_component(
server_args, component_name
)
component = self.load_customized(
component_model_path, server_args, component_name, **load_kwargs
)
# configure layerwise to make enough VRAM headroom
component = self._maybe_configure_layerwise_after_startup_cpu_staging(
component, server_args, component_name, load_kwargs
)
component = self._load_customized_with_context(
component_model_path,
server_args,
component_name,
attn_backend,
component_attn_name,
)
source = "sgl-diffusion"
except Exception as e:
if self.should_raise_customized_load_error(server_args, component_name):
@@ -240,18 +277,14 @@ class ComponentLoader(ABC):
f"Error while loading customized {component_name}, falling back to native version"
)
# fallback to native version
with component_attn_backend_context_manager(
attn_backend, component_name=component_attn_name
):
component = self.load_native(
component_model_path,
server_args,
transformers_or_diffusers,
component_name,
)
should_offload = self.should_offload(server_args)
target_device = self.target_device(should_offload)
component = component.to(device=target_device)
component = self._load_native_with_context(
component_model_path,
server_args,
component_name,
transformers_or_diffusers,
attn_backend,
component_attn_name,
)
source = "native"
logger.warning(
"Native component %s: %s is loaded, performance may be sub-optimal",
@@ -14,6 +14,7 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
resolve_transformer_safetensors_to_load,
)
from sglang.multimodal_gen.runtime.loader.utils import _normalize_component_type
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
@@ -155,12 +156,19 @@ class TransformerLoader(ComponentLoader):
else:
logger.debug("quantization config: %s", init_params["quant_config"])
local_torch_device = get_local_torch_device()
weight_load_plan = WeightLoadPlan.for_component(
checkpoint_load_device=local_torch_device,
needs_device_weight_postprocess=quant_spec.needs_device_weight_postprocess,
component_cpu_offload=bool(component_server_args.dit_cpu_offload),
)
# Load the model using FSDP loader
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params=init_params,
weight_dir_list=safetensors_list,
device=get_local_torch_device(),
device=local_torch_device,
hsdp_replicate_dim=server_args.hsdp_replicate_dim,
hsdp_shard_dim=server_args.hsdp_shard_dim,
cpu_offload=component_server_args.dit_cpu_offload,
@@ -170,10 +178,7 @@ class TransformerLoader(ComponentLoader):
reduce_dtype=torch.float32,
output_dtype=None,
strict=False,
defer_cpu_offload_until_after_weight_processing=(
component_server_args.dit_cpu_offload
and quant_spec.requires_device_weight_processing
),
weight_load_plan=weight_load_plan,
)
# post-hooks (e.g., patch scales (nunchaku))
@@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.loader.utils import (
hf_to_custom_state_dict,
set_default_torch_dtype,
)
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
from sglang.multimodal_gen.runtime.loader.weight_utils import (
safetensors_weights_iterator,
)
@@ -196,7 +197,7 @@ def maybe_load_fsdp_model(
output_dtype: torch.dtype | None = None,
pin_cpu_memory: bool = True,
strict: bool = True,
defer_cpu_offload_until_after_weight_processing: bool = False,
weight_load_plan: WeightLoadPlan | None = None,
) -> torch.nn.Module:
"""Load a model with optional FSDP (Fully Sharded Data Parallel) support.
@@ -207,12 +208,12 @@ def maybe_load_fsdp_model(
- Weight loading and casting
reduce_dtype: Data type for gradient reduction in FSDP mixed precision.
strict: If True, enforce strict state dict loading (all keys must match).
defer_cpu_offload_until_after_weight_processing: If True, keep weights
on device until process_weights_after_loading completes, then apply
non-FSDP CPU offload.
weight_load_plan: Optional checkpoint/postprocess device plan for this load.
"""
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
# manually casting the inputs to the model
# 1. prepare for loading
default_torch_dtype = param_dtype if param_dtype else torch.bfloat16
mp_policy = MixedPrecisionPolicy(
default_torch_dtype, reduce_dtype, output_dtype, cast_forward_inputs=False
@@ -236,8 +237,9 @@ def maybe_load_fsdp_model(
use_fsdp = False
logger.info("Disabling FSDP for MPS platform as it's not compatible")
weight_load_plan = weight_load_plan or WeightLoadPlan(checkpoint_load_device=device)
defer_cpu_offload = bool(
cpu_offload and defer_cpu_offload_until_after_weight_processing
cpu_offload and weight_load_plan.defer_component_cpu_offload
)
if defer_cpu_offload and use_fsdp:
logger.warning(
@@ -245,7 +247,11 @@ def maybe_load_fsdp_model(
"FSDP offload policy."
)
defer_cpu_offload = False
load_cpu_offload = cpu_offload and not defer_cpu_offload
load_cpu_offload = bool(cpu_offload and not defer_cpu_offload)
weight_postprocess_device = weight_load_plan.weight_postprocess_device
if use_fsdp and weight_postprocess_device is not None:
logger.warning("Ignoring weight postprocess device override for FSDP loading.")
weight_postprocess_device = None
if use_fsdp:
model._pre_fsdp_weight_loader_params = {
@@ -275,6 +281,8 @@ def maybe_load_fsdp_model(
)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
# 2. load model from disk
weight_iterator = safetensors_weights_iterator(weight_dir_list)
preprocess_loaded_state_dict = getattr(model, "preprocess_loaded_state_dict", None)
if preprocess_loaded_state_dict is not None:
@@ -292,7 +300,7 @@ def maybe_load_fsdp_model(
load_model_from_full_model_state_dict(
model,
weight_iterator,
device,
weight_load_plan.checkpoint_load_device,
param_dtype,
strict=strict,
cpu_offload=load_cpu_offload,
@@ -303,6 +311,11 @@ def maybe_load_fsdp_model(
dict(model.named_parameters()), bnb_quant_states
)
# 3. postprocessing
if weight_postprocess_device is not None:
# move to device to perform postprocessing
model.to(weight_postprocess_device)
for _, module in model.named_modules():
quant_method = getattr(module, "quant_method", None)
if quant_method is not None and hasattr(
@@ -316,8 +329,6 @@ def maybe_load_fsdp_model(
if _is_npu:
torch.npu.empty_cache()
model.post_load_weights()
if defer_cpu_offload:
model.to("cpu")
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
@@ -325,6 +336,11 @@ def maybe_load_fsdp_model(
# Avoid unintended computation graph accumulation during inference
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
# 4. deferred cpu offload
if defer_cpu_offload:
model.to("cpu")
return model
@@ -406,7 +422,7 @@ def shard_model(
def load_model_from_full_model_state_dict(
model: FSDPModule | torch.nn.Module,
full_sd_iterator: Generator[tuple[str, torch.Tensor], None, None],
device: torch.device,
checkpoint_load_device: torch.device,
param_dtype: torch.dtype | None,
strict: bool = False,
cpu_offload: bool = False,
@@ -418,7 +434,7 @@ def load_model_from_full_model_state_dict(
Args:
model (Union[FSDPModule, torch.nn.Module]): Model to generate fully qualified names for cpu_state_dict
full_sd_iterator (Generator): an iterator yielding (param_name, tensor) pairs
device (torch.device): device used to move full state dict tensors
checkpoint_load_device (torch.device): device used to move full state dict tensors
param_dtype (torch.dtype): dtype used to move full state dict tensors. If none, respect original dtype from checkpoint
strict (bool): flag to check if to load the model in strict mode
cpu_offload (bool): flag to check if FSDP offload is enabled
@@ -513,7 +529,9 @@ def load_model_from_full_model_state_dict(
)
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
full_tensor = full_tensor.to(
device=checkpoint_load_device, dtype=target_dtype
)
actual_param = _get_param_for_weight_loading(
model, param_dict, target_param_name
)
@@ -525,7 +543,9 @@ 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(
meta_sharded_param, device=device, dtype=target_dtype
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)
@@ -562,7 +582,9 @@ def load_model_from_full_model_state_dict(
if cpu_offload and not is_fsdp_model:
sharded_tensor = sharded_tensor.cpu()
else:
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
full_tensor = full_tensor.to(
device=checkpoint_load_device, dtype=target_dtype
)
actual_param = _get_param_for_weight_loading(
model, param_dict, target_param_name
)
@@ -575,7 +597,7 @@ def load_model_from_full_model_state_dict(
assert actual_param is not None
tp_sharded_tensor = torch.empty(
tuple(actual_param.shape),
device=device,
device=checkpoint_load_device,
dtype=target_dtype,
)
temp_param = _make_param_like(actual_param, tp_sharded_tensor)
@@ -722,13 +744,17 @@ def load_model_from_full_model_state_dict(
if not hasattr(meta_sharded_param, "device_mesh"):
sharded_tensor = init_like(
meta_sharded_param, device=device, dtype=meta_sharded_param_dtype
meta_sharded_param,
device=checkpoint_load_device,
dtype=meta_sharded_param_dtype,
)
if cpu_offload and not is_fsdp_model:
sharded_tensor = sharded_tensor.cpu()
else:
full_tensor = init_like(
meta_sharded_param, device=device, dtype=meta_sharded_param_dtype
meta_sharded_param,
device=checkpoint_load_device,
dtype=meta_sharded_param_dtype,
)
sharded_tensor = distribute_tensor(
full_tensor,
@@ -122,7 +122,7 @@ class TransformerQuantLoadSpec:
quant_config: Optional[QuantizationConfig]
nunchaku_config: Optional[NunchakuConfig]
param_dtype: Optional[torch.dtype]
requires_device_weight_processing: bool = False
needs_device_weight_postprocess: bool = False
post_load_hooks: list[PostLoadHook] = field(default_factory=list)
@property
@@ -482,23 +482,26 @@ def resolve_transformer_quant_load_spec(
quant_config=quant_config,
nunchaku_config=nunchaku_config,
param_dtype=param_dtype,
requires_device_weight_processing=_requires_device_weight_processing(
quant_config
),
needs_device_weight_postprocess=_needs_device_weight_postprocess(quant_config),
post_load_hooks=post_load_hooks,
)
def _requires_device_weight_processing(
def _needs_device_weight_postprocess(
quant_config: Optional[QuantizationConfig],
) -> bool:
"""Return whether post-load weight processing needs CUDA/NPU tensors."""
quant_name = _get_quant_config_name(quant_config)
if quant_name == "fp8":
return not getattr(quant_config, "is_checkpoint_fp8_serialized", False)
if quant_name == "mxfp4":
return not getattr(quant_config, "is_checkpoint_mxfp4_serialized", False)
return False
serialized_flag_by_quant_name = {
"fp8": "is_checkpoint_fp8_serialized",
"mxfp8": "is_checkpoint_fp8_serialized",
"mxfp4": "is_checkpoint_mxfp4_serialized",
"mxfp4_npu": "is_checkpoint_mxfp4_npu_serialized",
}
serialized_flag = serialized_flag_by_quant_name.get(quant_name)
if serialized_flag is None:
return False
return not getattr(quant_config, serialized_flag, False)
def _build_transformer_quant_adapters(
@@ -0,0 +1,35 @@
from dataclasses import dataclass
import torch
@dataclass(frozen=True)
class WeightLoadPlan:
"""Device plan for checkpoint loading, before runtime residency takes over."""
# Device used while materializing checkpoint tensors from files.
checkpoint_load_device: torch.device
# Device required while running process_weights_after_loading; None means unchanged.
weight_postprocess_device: torch.device | None = None
# Delay non-FSDP component CPU offload until after weight postprocessing.
defer_component_cpu_offload: bool = False
@classmethod
def for_component(
cls,
*,
checkpoint_load_device: torch.device,
needs_device_weight_postprocess: bool,
component_cpu_offload: bool,
) -> "WeightLoadPlan":
# if on-device weight postprocessing is required, load directly to device to speedup loading
weight_postprocess_device = (
checkpoint_load_device if needs_device_weight_postprocess else None
)
return cls(
checkpoint_load_device=checkpoint_load_device,
weight_postprocess_device=weight_postprocess_device,
defer_component_cpu_offload=(
needs_device_weight_postprocess and component_cpu_offload
),
)
@@ -57,10 +57,11 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
_filter_duplicate_precision_variant_safetensors,
_Flux2Nvfp4FallbackAdapter,
_requires_device_weight_processing,
_needs_device_weight_postprocess,
resolve_transformer_quant_load_spec,
resolve_transformer_safetensors_to_load,
)
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
from sglang.multimodal_gen.runtime.models.dits.flux import FluxSingleTransformerBlock
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
build_nvfp4_config_from_safetensors_list,
@@ -81,6 +82,18 @@ class _FakeQuantConfig:
return "modelopt_fp4"
def _make_quant_config(name: str, **attrs):
cls = type(
f"_Fake{name.title().replace('_', '')}QuantConfig",
(),
{"get_name": classmethod(lambda cls: name)},
)
quant_config = cls()
for attr_name, attr_value in attrs.items():
setattr(quant_config, attr_name, attr_value)
return quant_config
class TestTransformerQuantHelpers(unittest.TestCase):
def _make_server_args(self, **overrides):
defaults = dict(
@@ -184,13 +197,40 @@ class TestTransformerQuantHelpers(unittest.TestCase):
self.assertEqual(resolved, files)
def test_online_fp8_requires_device_weight_processing(self):
self.assertTrue(_requires_device_weight_processing(Fp8Config()))
def test_weight_load_plan_defers_cpu_offload_for_device_postprocess(self):
device = torch.device("cuda:0")
plan = WeightLoadPlan.for_component(
checkpoint_load_device=device,
needs_device_weight_postprocess=True,
component_cpu_offload=True,
)
self.assertEqual(plan.checkpoint_load_device, device)
self.assertEqual(plan.weight_postprocess_device, device)
self.assertTrue(plan.defer_component_cpu_offload)
def test_online_fp8_needs_device_weight_postprocess(self):
self.assertTrue(_needs_device_weight_postprocess(Fp8Config()))
self.assertFalse(
_requires_device_weight_processing(
_needs_device_weight_postprocess(
Fp8Config(is_checkpoint_fp8_serialized=True)
)
)
self.assertTrue(_needs_device_weight_postprocess(_make_quant_config("mxfp8")))
self.assertFalse(
_needs_device_weight_postprocess(
_make_quant_config("mxfp8", is_checkpoint_fp8_serialized=True)
)
)
self.assertTrue(
_needs_device_weight_postprocess(_make_quant_config("mxfp4_npu"))
)
self.assertFalse(
_needs_device_weight_postprocess(
_make_quant_config("mxfp4_npu", is_checkpoint_mxfp4_npu_serialized=True)
)
)
@patch(
"sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list",