[diffusion] feat: add memory-aware component load order (#25457)
This commit is contained in:
+174
@@ -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,
|
||||||
|
),
|
||||||
|
)
|
||||||
@@ -12,6 +12,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_s
|
|||||||
LayerwiseOffloadStrategy,
|
LayerwiseOffloadStrategy,
|
||||||
ResidentStrategy,
|
ResidentStrategy,
|
||||||
VanillaD2HStrategy,
|
VanillaD2HStrategy,
|
||||||
|
is_fsdp_managed_module,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
is_layerwise_offloaded_module,
|
is_layerwise_offloaded_module,
|
||||||
@@ -93,10 +94,6 @@ class ComponentResidencyPipeline(Protocol):
|
|||||||
component_residency_strategies: MutableMapping[str, "ComponentResidencyStrategy"]
|
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(
|
def should_cpu_offload_component(
|
||||||
component_name: str, module: nn.Module, server_args: ServerArgs
|
component_name: str, module: nn.Module, server_args: ServerArgs
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
|||||||
+6
-2
@@ -57,6 +57,10 @@ def _module_ready_on_local_device(
|
|||||||
return dtype is None or tensor.dtype == dtype
|
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:
|
class ComponentResidencyStrategy:
|
||||||
"""Baseclass for describing how a component should be treated (regarding where its weights locates)
|
"""Baseclass for describing how a component should be treated (regarding where its weights locates)
|
||||||
|
|
||||||
@@ -75,7 +79,6 @@ class ComponentResidencyStrategy:
|
|||||||
use: ComponentUse,
|
use: ComponentUse,
|
||||||
state: ResidencyState,
|
state: ResidencyState,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""hook called"""
|
|
||||||
self.enter(module)
|
self.enter(module)
|
||||||
|
|
||||||
def wait_for_use(
|
def wait_for_use(
|
||||||
@@ -144,7 +147,8 @@ class ResidentStrategy(ComponentResidencyStrategy):
|
|||||||
use: ComponentUse,
|
use: ComponentUse,
|
||||||
state: ResidencyState,
|
state: ResidencyState,
|
||||||
) -> None:
|
) -> None:
|
||||||
if use.target_dtype is not None:
|
if is_fsdp_managed_module(module):
|
||||||
|
return
|
||||||
_module_to_local_device(module, dtype=use.target_dtype)
|
_module_to_local_device(module, dtype=use.target_dtype)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -24,6 +24,10 @@ from sglang.multimodal_gen.runtime.layers.attention.selector import (
|
|||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
PipelineComponentLoader,
|
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 (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||||
ComponentResidencyManager,
|
ComponentResidencyManager,
|
||||||
ComponentResidencyStrategy,
|
ComponentResidencyStrategy,
|
||||||
@@ -385,10 +389,16 @@ class ComposedPipelineBase(ABC):
|
|||||||
logger.info("Loading required components: %s", required_modules)
|
logger.info("Loading required components: %s", required_modules)
|
||||||
|
|
||||||
loaded_components = {}
|
loaded_components = {}
|
||||||
for module_name, (
|
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,
|
transformers_or_diffusers,
|
||||||
architecture,
|
architecture,
|
||||||
) in tqdm(iterable=model_index.items(), desc="Loading required modules"):
|
),
|
||||||
|
) in enumerate(model_index.items()):
|
||||||
if transformers_or_diffusers is None:
|
if transformers_or_diffusers is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Module %s in model_index.json has null value, removing from required_config_modules",
|
"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]
|
loaded_components[module_name] = loaded_modules[module_name]
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# we load the module from the extra config module map if it exists
|
|
||||||
if module_name in self._extra_config_module_map:
|
if module_name in self._extra_config_module_map:
|
||||||
load_module_name = self._extra_config_module_map[module_name]
|
load_module_name = self._extra_config_module_map[module_name]
|
||||||
else:
|
else:
|
||||||
load_module_name = module_name
|
load_module_name = module_name
|
||||||
|
|
||||||
component_model_path = self._resolve_component_path(
|
component_model_path = self._resolve_component_path(
|
||||||
server_args, module_name, load_module_name
|
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 = (
|
attn_backend, matched_backend_key = (
|
||||||
server_args.resolve_component_attention_backend(
|
server_args.resolve_component_attention_backend(
|
||||||
module_name, load_module_name
|
module_name, load_module_name
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -9,10 +9,14 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
|||||||
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
||||||
_ModelOptFp8OffloadAdapter,
|
_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 (
|
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
||||||
layerwise_offload as layerwise_offload_mod,
|
layerwise_offload as layerwise_offload_mod,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||||
|
ComponentUse,
|
||||||
build_component_residency_strategy,
|
build_component_residency_strategy,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_strategies import (
|
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)
|
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):
|
def test_layerwise_offload_aligns_contiguous_tensor_offsets(monkeypatch):
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
|
layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule
|
||||||
|
|||||||
Reference in New Issue
Block a user