diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_loading_order.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_loading_order.py new file mode 100644 index 000000000..890eaf92e --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_loading_order.py @@ -0,0 +1,174 @@ +"""Memory-aware ordering for pipeline component weight loads to avoid OOM while loading. + +Load the VRAM-intensive components earlier than others + +The pipeline owns component selection, path resolution, and actual loading; this +module only ranks already-selected load specs. +""" + +import glob +import json +import os +from dataclasses import dataclass + +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( + is_dit_component_name, + is_image_encoder_component_name, + is_text_encoder_component_name, + is_vae_component_name, +) + + +@dataclass(frozen=True) +class ComponentLoadSpec: + """One pipeline component that still needs a real weight load.""" + + module_name: str + load_module_name: str + component_model_path: str + transformers_or_diffusers: str + architecture: str | None + index: int + + +_WEIGHT_FILE_SUFFIXES = (".bin", ".pt", ".pth") + + +def _component_base_name(component_name: str) -> str: + prefix, separator, suffix = component_name.rpartition("_") + if separator and suffix.isdigit(): + return prefix + return component_name + + +def _component_variant_priority(component_name: str) -> int: + _, separator, suffix = component_name.rpartition("_") + if separator and suffix.isdigit(): + return -int(suffix) + return 0 + + +def component_load_risk_rank(component_name: str) -> int: + """Fallback type rank when checkpoint size cannot be inferred.""" + candidate_names = (component_name, _component_base_name(component_name)) + if any(is_dit_component_name(name) for name in candidate_names): + return 0 + if any(is_text_encoder_component_name(name) for name in candidate_names): + return 1 + if any(is_image_encoder_component_name(name) for name in candidate_names): + return 2 + if any(is_vae_component_name(name) for name in candidate_names): + return 3 + return 10 + + +def _safe_file_size(file_path: str) -> int | None: + try: + return os.path.getsize(file_path) + except OSError: + return None + + +def _safetensors_payload_size_bytes(file_path: str) -> int | None: + try: + with open(file_path, "rb") as f: + header_size_bytes = f.read(8) + if len(header_size_bytes) != 8: + return _safe_file_size(file_path) + header_size = int.from_bytes(header_size_bytes, "little") + header = json.loads(f.read(header_size)) + except (OSError, json.JSONDecodeError, ValueError): + return _safe_file_size(file_path) + + payload_size = 0 + for tensor_name, tensor_info in header.items(): + if tensor_name == "__metadata__": + continue + offsets = tensor_info.get("data_offsets") + if not isinstance(offsets, list) or len(offsets) != 2: + return _safe_file_size(file_path) + payload_size += offsets[1] - offsets[0] + return payload_size + + +def _safetensors_files_from_index(component_model_path: str) -> list[str]: + indexed_files: set[str] = set() + index_paths = sorted( + glob.glob(os.path.join(component_model_path, "*.safetensors.index.json")) + ) + for index_path in index_paths: + try: + with open(index_path) as f: + weight_map = json.load(f).get("weight_map", {}) + except (OSError, json.JSONDecodeError): + continue + for shard_name in weight_map.values(): + shard_path = os.path.join(component_model_path, shard_name) + if os.path.isfile(shard_path): + indexed_files.add(shard_path) + return sorted(indexed_files) + + +def _list_component_safetensors_files(component_model_path: str) -> list[str]: + if os.path.isfile(component_model_path): + if component_model_path.endswith(".safetensors"): + return [component_model_path] + return [] + if not os.path.isdir(component_model_path): + return [] + + indexed_files = _safetensors_files_from_index(component_model_path) + if indexed_files: + return indexed_files + return sorted(glob.glob(os.path.join(component_model_path, "*.safetensors"))) + + +def infer_component_weight_size_bytes(component_model_path: str) -> int | None: + """Infer checkpoint payload size from safetensors without materializing tensors.""" + safetensors_files = _list_component_safetensors_files(component_model_path) + if safetensors_files: + sizes = [ + size + for size in ( + _safetensors_payload_size_bytes(file_path) + for file_path in safetensors_files + ) + if size is not None + ] + return sum(sizes) if sizes else None + + if os.path.isfile(component_model_path): + if component_model_path.endswith(_WEIGHT_FILE_SUFFIXES): + return _safe_file_size(component_model_path) + return None + if not os.path.isdir(component_model_path): + return None + + weight_files = [] + for suffix in _WEIGHT_FILE_SUFFIXES: + weight_files.extend(glob.glob(os.path.join(component_model_path, f"*{suffix}"))) + if not weight_files: + return None + sizes = [ + size + for size in (_safe_file_size(file_path) for file_path in weight_files) + if size is not None + ] + return sum(sizes) if sizes else None + + +def order_component_load_specs( + component_specs: list[ComponentLoadSpec], +) -> list[ComponentLoadSpec]: + # load larger weight payloads before small helpers to reduce startup peak OOMs + return sorted( + component_specs, + key=lambda spec: ( + # 1. model size inferred from checkpoints + -(infer_component_weight_size_bytes(spec.component_model_path) or 0), + # 2. infer from component name + component_load_risk_rank(spec.load_module_name), + _component_variant_priority(spec.load_module_name), + spec.index, + ), + ) diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py index 256214438..7501257d9 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py @@ -12,6 +12,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_s LayerwiseOffloadStrategy, ResidentStrategy, VanillaD2HStrategy, + is_fsdp_managed_module, ) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( is_layerwise_offloaded_module, @@ -93,10 +94,6 @@ class ComponentResidencyPipeline(Protocol): component_residency_strategies: MutableMapping[str, "ComponentResidencyStrategy"] -def is_fsdp_managed_module(module: nn.Module) -> bool: - return module.__class__.__name__.startswith("FSDP") - - def should_cpu_offload_component( component_name: str, module: nn.Module, server_args: ServerArgs ) -> bool: diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_resident_strategies.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_resident_strategies.py index eef4c9c5d..bd51cda1c 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_resident_strategies.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_resident_strategies.py @@ -57,6 +57,10 @@ def _module_ready_on_local_device( return dtype is None or tensor.dtype == dtype +def is_fsdp_managed_module(module: nn.Module) -> bool: + return module.__class__.__name__.startswith("FSDP") + + class ComponentResidencyStrategy: """Baseclass for describing how a component should be treated (regarding where its weights locates) @@ -75,7 +79,6 @@ class ComponentResidencyStrategy: use: ComponentUse, state: ResidencyState, ) -> None: - """hook called""" self.enter(module) def wait_for_use( @@ -144,8 +147,9 @@ class ResidentStrategy(ComponentResidencyStrategy): use: ComponentUse, state: ResidencyState, ) -> None: - if use.target_dtype is not None: - _module_to_local_device(module, dtype=use.target_dtype) + if is_fsdp_managed_module(module): + return + _module_to_local_device(module, dtype=use.target_dtype) class SnapshotModuleResidency: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py index 316896637..bb3639f69 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py @@ -24,6 +24,10 @@ from sglang.multimodal_gen.runtime.layers.attention.selector import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( PipelineComponentLoader, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.component_loading_order import ( + ComponentLoadSpec, + order_component_load_specs, +) from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( ComponentResidencyManager, ComponentResidencyStrategy, @@ -385,10 +389,16 @@ class ComposedPipelineBase(ABC): logger.info("Loading required components: %s", required_modules) loaded_components = {} - for module_name, ( - transformers_or_diffusers, - architecture, - ) in tqdm(iterable=model_index.items(), desc="Loading required modules"): + component_load_specs: list[ComponentLoadSpec] = [] + + # enqueue only real weight loads (e.g., scheduler, tokenizer is excluded); skipped/provided modules keep old handling + for index, ( + module_name, + ( + transformers_or_diffusers, + architecture, + ), + ) in enumerate(model_index.items()): if transformers_or_diffusers is None: logger.warning( "Module %s in model_index.json has null value, removing from required_config_modules", @@ -405,15 +415,43 @@ class ComposedPipelineBase(ABC): loaded_components[module_name] = loaded_modules[module_name] continue - # we load the module from the extra config module map if it exists if module_name in self._extra_config_module_map: load_module_name = self._extra_config_module_map[module_name] else: load_module_name = module_name - component_model_path = self._resolve_component_path( server_args, module_name, load_module_name ) + # collect loading specs + component_load_specs.append( + ComponentLoadSpec( + module_name=module_name, + load_module_name=load_module_name, + component_model_path=component_model_path, + transformers_or_diffusers=transformers_or_diffusers, + architecture=architecture, + index=index, + ) + ) + + # reorder loading order to avoid OOM + component_load_specs: ComponentLoadSpec = order_component_load_specs( + component_load_specs + ) + logger.info( + "Memory-aware component load order: %s", + [spec.module_name for spec in component_load_specs], + ) + + for spec in tqdm( + iterable=component_load_specs, desc="Loading required modules" + ): + module_name: str = spec.module_name + load_module_name: str = spec.load_module_name + transformers_or_diffusers: str = spec.transformers_or_diffusers + architecture: str = spec.architecture + component_model_path: str = spec.component_model_path + attn_backend, matched_backend_key = ( server_args.resolve_component_attention_backend( module_name, load_module_name diff --git a/python/sglang/multimodal_gen/test/unit/test_component_loading_order.py b/python/sglang/multimodal_gen/test/unit/test_component_loading_order.py new file mode 100644 index 000000000..d0f0f7294 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_component_loading_order.py @@ -0,0 +1,130 @@ +import json + +from sglang.multimodal_gen.runtime.managers.memory_managers.component_loading_order import ( + ComponentLoadSpec, + component_load_risk_rank, + infer_component_weight_size_bytes, + order_component_load_specs, +) + + +def _spec( + component_name: str, index: int, component_model_path: str = "/missing" +) -> ComponentLoadSpec: + return ComponentLoadSpec( + module_name=component_name, + load_module_name=component_name, + component_model_path=component_model_path, + transformers_or_diffusers="diffusers", + architecture=None, + index=index, + ) + + +def _write_safetensors(path, payload_size: int) -> None: + header = { + "weight": { + "dtype": "F16", + "shape": [payload_size // 2], + "data_offsets": [0, payload_size], + } + } + header_bytes = json.dumps(header).encode("utf-8") + path.write_bytes( + len(header_bytes).to_bytes(8, "little") + header_bytes + b"\0" * payload_size + ) + + +def test_component_load_order_prioritizes_weight_heavy_components(): + specs = [ + _spec("scheduler", 0), + _spec("tokenizer", 1), + _spec("text_encoder", 2), + _spec("transformer", 3), + _spec("vae", 4), + ] + + ordered_names = [spec.module_name for spec in order_component_load_specs(specs)] + + assert ordered_names == [ + "transformer", + "text_encoder", + "vae", + "scheduler", + "tokenizer", + ] + + +def test_component_load_order_prioritizes_larger_numbered_variants(): + specs = [ + _spec("transformer", 0), + _spec("transformer_2", 1), + _spec("text_encoder", 2), + _spec("text_encoder_3", 3), + _spec("text_encoder_2", 4), + ] + + ordered_names = [spec.module_name for spec in order_component_load_specs(specs)] + + assert ordered_names == [ + "transformer_2", + "transformer", + "text_encoder_3", + "text_encoder_2", + "text_encoder", + ] + + +def test_component_load_order_uses_load_module_name_for_extra_config_alias(): + specs = [ + ComponentLoadSpec( + module_name="condition_image_encoder", + load_module_name="condition_image_encoder", + component_model_path="/missing", + transformers_or_diffusers="diffusers", + architecture=None, + index=0, + ), + ComponentLoadSpec( + module_name="encoder_alias", + load_module_name="text_encoder_2", + component_model_path="/missing", + transformers_or_diffusers="transformers", + architecture=None, + index=1, + ), + ] + + ordered_names = [spec.module_name for spec in order_component_load_specs(specs)] + + assert ordered_names == ["encoder_alias", "condition_image_encoder"] + + +def test_component_load_risk_rank_keeps_small_helpers_last(): + assert component_load_risk_rank("transformer") < component_load_risk_rank( + "scheduler" + ) + assert component_load_risk_rank("text_encoder_2") < component_load_risk_rank( + "processor" + ) + assert component_load_risk_rank("vae") < component_load_risk_rank("tokenizer") + + +def test_component_load_order_prefers_inferred_safetensors_size(tmp_path): + small_transformer_path = tmp_path / "small_transformer" + large_encoder_path = tmp_path / "large_encoder" + small_transformer_path.mkdir() + large_encoder_path.mkdir() + _write_safetensors(small_transformer_path / "model.safetensors", 16) + _write_safetensors(large_encoder_path / "model.safetensors", 64) + + specs = [ + _spec("transformer", 0, str(small_transformer_path)), + _spec("text_encoder", 1, str(large_encoder_path)), + _spec("scheduler", 2), + ] + + ordered_names = [spec.module_name for spec in order_component_load_specs(specs)] + + assert ordered_names == ["text_encoder", "transformer", "scheduler"] + assert infer_component_weight_size_bytes(str(large_encoder_path)) == 64 diff --git a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py index e9d49a33b..22c62311f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -9,10 +9,14 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( _ModelOptFp8OffloadAdapter, ) +from sglang.multimodal_gen.runtime.managers.memory_managers import ( + component_resident_strategies as component_resident_strategies_mod, +) from sglang.multimodal_gen.runtime.managers.memory_managers import ( layerwise_offload as layerwise_offload_mod, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + ComponentUse, build_component_residency_strategy, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_strategies import ( @@ -346,6 +350,50 @@ def test_component_cpu_offload_strategy_remains_flag_driven(): assert isinstance(strategy, ResidentStrategy) +def test_resident_strategy_prepares_local_device_without_dtype(monkeypatch): + calls = [] + + def fake_module_to_local_device(module, *, dtype=None): + calls.append((module, dtype)) + + monkeypatch.setattr( + component_resident_strategies_mod, + "_module_to_local_device", + fake_module_to_local_device, + ) + module = _DummyModel() + + ResidentStrategy().prepare_for_use( + module, + ComponentUse(stage_name="DenoisingStage", component_name="transformer"), + SimpleNamespace(), + ) + + assert calls == [(module, None)] + + +def test_resident_strategy_keeps_fsdp_managed_module_owned_by_fsdp(monkeypatch): + calls = [] + + def fake_module_to_local_device(module, *, dtype=None): + calls.append((module, dtype)) + + monkeypatch.setattr( + component_resident_strategies_mod, + "_module_to_local_device", + fake_module_to_local_device, + ) + module = type("FSDPDummyModel", (_DummyModel,), {})() + + ResidentStrategy().prepare_for_use( + module, + ComponentUse(stage_name="TextEncodingStage", component_name="text_encoder"), + SimpleNamespace(), + ) + + assert calls == [] + + def test_layerwise_offload_aligns_contiguous_tensor_offsets(monkeypatch): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule