diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index a3d4a08f3..d67500b72 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -23,6 +23,7 @@ if TYPE_CHECKING: SGLANG_DIFFUSION_TRACE_FUNCTION: int = 0 SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST: bool = False SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE: bool = False + SGLANG_DIFFUSION_DISABLE_LORA_MERGE_CACHE: bool = False SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork" SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda" SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = "" @@ -278,6 +279,11 @@ environment_variables: dict[str, Callable[[], Any]] = { "SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE": _lazy_bool( "SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE" ), + # Kill-switch: keep LoRA-merged weights in anonymous host memory instead + # of the file-backed LoRA merge cache. + "SGLANG_DIFFUSION_DISABLE_LORA_MERGE_CACHE": _lazy_bool( + "SGLANG_DIFFUSION_DISABLE_LORA_MERGE_CACHE" + ), # ================== cache-dit Env Vars ================== # Enable cache-dit acceleration for DiT inference # CUDA-IPC transport for 2-rank Ulysses all-to-all (NVLink same-node) diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index 272331fce..542144779 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -74,14 +74,24 @@ class BaseLayerWithLoRA(nn.Module): base_layer: nn.Module, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ): super().__init__() self.base_layer: nn.Module = base_layer self.merged: bool = False # Immutable base-weight snapshot; `to("cpu")` may alias CPU storage. - # Use `clone()` so merge updates cannot mutate this backup tensor. - self.cpu_weight = base_layer.weight.detach().to("cpu").clone() + # Use `clone()` so in-place merge updates cannot mutate this backup. + # With snapshot_base=False the snapshot is a zero-copy view instead: + # valid only while every merge on this layer is a copy-merge (the + # merged-store path), which never writes the base storage. H3's DiT + # backup alone is 38 GB of anonymous memory under clone(). + if snapshot_base: + self.cpu_weight = base_layer.weight.detach().to("cpu").clone() + self._base_is_view = False + else: + self.cpu_weight = base_layer.weight.detach() + self._base_is_view = True # indicates adapter weights don't contain this layer # (which shouldn't normally happen, but we want to separate it from the case of erroneous merging) # Default to True to prevent using uninitialized weights; set to False when weights are loaded @@ -205,6 +215,13 @@ class BaseLayerWithLoRA(nn.Module): elif self.merged: self.unmerge_lora_weights() + def _ensure_base_snapshot_owned(self) -> None: + """An in-place merge is about to write the base storage; if the + snapshot is a zero-copy view into it, materialize the clone now.""" + if self._base_is_view: + self.cpu_weight = self.cpu_weight.clone() + self._base_is_view = False + @torch.no_grad() def _merge_lora_into_data( self, @@ -274,6 +291,41 @@ class BaseLayerWithLoRA(nn.Module): return False return True + @torch.no_grad() + def compute_merged_weight(self) -> torch.Tensor: + """The merged weight as a new CPU tensor; the base is never written. + + Same math as the in-place merge — computed on the device, in fp32 + when the policy says so, rounded back once — so the bytes are + identical to what merge_lora_weights would have left in place. + """ + base = self.weight.data + target_dtype = base.dtype + work = base.detach().to(get_local_torch_device()) + if ( + self._should_merge_in_fp32(self.lora_weights_list) + and work.is_floating_point() + and work.dtype != torch.float32 + ): + work = work.to(torch.float32) + self._merge_lora_into_data(work, self.lora_weights_list) + return work.to("cpu", dtype=target_dtype) + + def install_merged_weight( + self, merged: torch.Tensor, base_view: torch.Tensor + ) -> None: + """Adopt an externally held merged weight (e.g. a cache mapping). + + The single place the cached-merge state transition happens: the + parameter points at `merged`, the layer counts as merged, and the + unmerge snapshot is the untouched base view — zero-copy, because + nothing wrote the base storage. + """ + self.weight.data = merged + self.merged = True + self.cpu_weight = base_view.detach() + self._base_is_view = True + @torch.no_grad() def merge_lora_weights(self, strength: float | None = None) -> None: if strength is not None: @@ -294,6 +346,7 @@ class BaseLayerWithLoRA(nn.Module): if self.disable_lora: return + self._ensure_base_snapshot_owned() if self.merged: self.unmerge_lora_weights() @@ -476,8 +529,9 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): base_layer: ColumnParallelLinear, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ) -> None: - super().__init__(base_layer, lora_rank, lora_alpha) + super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) def forward(self, input_: torch.Tensor) -> torch.Tensor: if self.merged or self.disable_lora: @@ -538,8 +592,9 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA): base_layer: MergedColumnParallelLinear, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ) -> None: - super().__init__(base_layer, lora_rank, lora_alpha) + super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor: return A @@ -574,8 +629,9 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA): base_layer: QKVParallelLinear, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ) -> None: - super().__init__(base_layer, lora_rank, lora_alpha) + super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor: return A @@ -606,8 +662,9 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): base_layer: RowParallelLinear, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ) -> None: - super().__init__(base_layer, lora_rank, lora_alpha) + super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) def forward(self, input_: torch.Tensor): if self.merged or self.disable_lora: @@ -692,8 +749,9 @@ class LinearWithLoRA(BaseLayerWithLoRA): base_layer: nn.Linear, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ) -> None: - super().__init__(base_layer, lora_rank, lora_alpha) + super().__init__(base_layer, lora_rank, lora_alpha, snapshot_base) @torch.compile() def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -728,10 +786,15 @@ class LinearWithLoRA(BaseLayerWithLoRA): return out +def _use_owned_base_snapshot(snapshot_base: bool, device_type: str) -> bool: + return snapshot_base or device_type not in ("cpu", "meta") + + def wrap_with_lora_layer( layer: nn.Module, lora_rank: int | None = None, lora_alpha: int | None = None, + snapshot_base: bool = True, ) -> BaseLayerWithLoRA | None: """ transform the given layer to its corresponding LoRA layer @@ -750,10 +813,14 @@ def wrap_with_lora_layer( } for src_layer_type, lora_layer_type in supported_layer_types.items(): if isinstance(layer, src_layer_type): # type: ignore[arg-type] + effective_snapshot_base = _use_owned_base_snapshot( + snapshot_base, layer.weight.device.type + ) ret = lora_layer_type( layer, lora_rank=lora_rank, lora_alpha=lora_alpha, + snapshot_base=effective_snapshot_base, ) return ret return None diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/lora_merge_cache.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/lora_merge_cache.py new file mode 100644 index 000000000..f7de9422b --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/lora_merge_cache.py @@ -0,0 +1,200 @@ +"""File-backed store for LoRA-merged weights. + +Merging an adapter writes the base weight in place. Under layerwise offload +the base weight is a view into the checkpoint mapping, so the write is a +copy-on-write: every merged byte turns into anonymous host memory the kernel +cannot reclaim. MiniMax-H3's DiT alone is 61.7 GB — a real 32 GB host dies on +it, and on any host the pin budget collapses to zero before the offload +managers ever see the weights. + +Written once to a per-layer cache file and mapped back, the same merged bytes +become page cache: droppable, refaultable, and invisible to the anonymous +accounting. The offload managers then classify them as mapped weights on +their own — no coordination needed. Rehoming happens layer by layer inside +the merge loop, so the anonymous high-water mark stays one layer wide, and a +later start with the same (base, adapters, strengths) adopts the store +without paying the merge at all. +""" + +import hashlib +import json +import os +import shutil + +import torch +from safetensors.torch import load_file as safetensors_load_file +from safetensors.torch import save_file as safetensors_save_file + +from sglang.multimodal_gen import envs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +_MANIFEST = "manifest.json" +_DISK_HEADROOM = 1.15 + + +def lora_merge_cache_key( + base_paths: list[str], + adapters: list[tuple[str, float, float | None]], +) -> str: + """Key of one merged-weights combination. + + `base_paths` are the component checkpoint paths (HF snapshot paths carry + the revision hash); `adapters` are ordered (lora_path, strength, alpha) + triples — order matters, merges compose in order. + """ + parts = [os.path.realpath(p) for p in sorted(base_paths)] + for path, strength, alpha in adapters: + real = os.path.realpath(path) + try: + size = os.path.getsize(real) + except OSError: + size = -1 + parts.append(f"{real}|{size}|{strength}|{alpha}") + return hashlib.sha1("||".join(parts).encode()).hexdigest()[:16] + + +class LoraMergeCache: + """Streams merged weights into a cache directory, one file per layer.""" + + def __init__(self, key: str, expected_bytes: int) -> None: + self.root = os.path.join( + envs.SGLANG_DIFFUSION_CACHE_ROOT, "lora_merge_cache", key + ) + self.manifest_path = os.path.join(self.root, _MANIFEST) + self.expected_bytes = expected_bytes + self._entries: dict[str, dict] = {} + self._writable: bool | None = None + + # -- adoption (fast path) ------------------------------------------------- + + def is_complete(self) -> bool: + """A complete store from an earlier run of the same combination.""" + try: + with open(self.manifest_path) as handle: + manifest = json.load(handle) + except (OSError, ValueError): + return False + entries = manifest.get("layers") + if not isinstance(entries, dict) or not entries: + return False + for meta in entries.values(): + if not os.path.exists(os.path.join(self.root, meta.get("file", ""))): + return False + self._entries = entries + return True + + def get( + self, name: str, shape: torch.Size, dtype: torch.dtype + ) -> torch.Tensor | None: + """The cached merged tensor for `name`, mapped from its file. + + Purely a lookup: the caller decides what to do with the tensor. A + missing or mismatched entry returns None — mismatch also drops the + remaining entries, because one wrong file means the whole combination + key no longer describes this module. + """ + meta = self._entries.get(name) + if meta is None: + return None + mapped = safetensors_load_file(os.path.join(self.root, meta["file"])) + tensor = mapped.get("weight") + if ( + tensor is None + or tuple(tensor.shape) != tuple(shape) + or tensor.dtype != dtype + ): + logger.warning( + "LoRA merge cache entry for %s does not match the module; " + "ignoring the cache", + name, + ) + self._entries = {} + return None + return tensor + + # -- capture (first run) -------------------------------------------------- + + def _ensure_writable(self) -> bool: + if self._writable is not None: + return self._writable + try: + os.makedirs(self.root, exist_ok=True) + usage = shutil.disk_usage(self.root) + if usage.free < self.expected_bytes * _DISK_HEADROOM: + logger.warning( + "LoRA merge cache needs %.1f GiB free under %s but only " + "%.1f GiB is available; merged weights stay in anonymous " + "host memory", + self.expected_bytes * _DISK_HEADROOM / 1024**3, + self.root, + usage.free / 1024**3, + ) + self._writable = False + else: + self._writable = True + except OSError as exc: + logger.warning("LoRA merge cache unavailable (%s)", exc) + self._writable = False + return self._writable + + def put(self, name: str, merged: torch.Tensor) -> torch.Tensor | None: + """Write one merged tensor to its cache file and return the mapping. + + The returned tensor is a view into the file — page cache the kernel + can drop — and the only thing the cache hands back; what to install it + into is the caller's business. None means the bytes could not be + cached (disk shortage, write failure) and the caller should keep its + own copy. + """ + if not self._ensure_writable(): + return None + fname = hashlib.sha1(name.encode()).hexdigest()[:16] + ".safetensors" + path = os.path.join(self.root, fname) + try: + tmp = f"{path}.tmp.{os.getpid()}" + safetensors_save_file({"weight": merged.contiguous()}, tmp) + os.replace(tmp, path) + mapped = safetensors_load_file(path)["weight"] + except Exception as exc: + logger.warning( + "Could not cache merged weight %s (%s); it stays in " + "anonymous host memory", + name, + exc, + ) + try: + if os.path.exists(path): + os.remove(path) + except OSError: + pass + return None + self._entries[name] = { + "file": fname, + "shape": list(merged.shape), + "dtype": str(merged.dtype), + } + return mapped + + def finalize(self, extra: dict | None = None) -> None: + """Write the manifest; only a complete store is ever adopted.""" + if not self._entries or not self._ensure_writable(): + return + manifest = {"layers": self._entries} + if extra: + manifest.update(extra) + tmp = f"{self.manifest_path}.tmp.{os.getpid()}" + try: + with open(tmp, "w") as handle: + json.dump(manifest, handle) + os.replace(tmp, self.manifest_path) + except OSError as exc: + logger.warning("LoRA merge cache manifest not written (%s)", exc) + return + logger.info( + "Merged weights cached to %s (%d layers); anonymous host memory " + "no longer holds them", + self.root, + len(self._entries), + ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py index eac9cd04d..a1afd5bec 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py @@ -12,6 +12,7 @@ import torch.distributed as dist from safetensors.torch import load_file from torch.distributed.tensor import DTensor +from sglang.multimodal_gen import envs from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.layers.lora.linear import ( BaseLayerWithLoRA, @@ -28,6 +29,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import from sglang.multimodal_gen.runtime.pipelines_core.lora.format_adapter import ( normalize_lora_state_dict, ) +from sglang.multimodal_gen.runtime.pipelines_core.lora.lora_merge_cache import ( + LoraMergeCache, + lora_merge_cache_key, +) from sglang.multimodal_gen.runtime.pipelines_core.lora.peft_adapter import ( get_peft_lora_alpha, load_peft_config, @@ -187,12 +192,13 @@ class LoRAPipeline(ComposedPipelineBase): self.lora_path = self.server_args.lora_path self.lora_nickname = self.server_args.lora_nickname if self.lora_path is not None: - self.convert_to_lora_layers() + self.convert_to_lora_layers(snapshot_base=False) self.set_lora( self.lora_nickname, self.lora_path, strength=self.server_args.lora_scale, # type: ignore lora_alpha=self.server_args.lora_alpha, + cache_merged=True, ) # type: ignore def is_target_layer(self, module_name: str) -> bool: @@ -323,6 +329,7 @@ class LoRAPipeline(ComposedPipelineBase): module_name: str, target_lora_layers: dict[str, BaseLayerWithLoRA], check_exclude: bool = True, + snapshot_base: bool = True, ) -> int: """ Convert layers in a module to LoRA layers. @@ -352,6 +359,7 @@ class LoRAPipeline(ComposedPipelineBase): layer, lora_rank=self.lora_rank, lora_alpha=self.lora_alpha, + snapshot_base=snapshot_base, ) if lora_layer is not None: target_lora_layers[name] = lora_layer @@ -385,9 +393,14 @@ class LoRAPipeline(ComposedPipelineBase): "unquantized checkpoint to use LoRA." ) - def convert_to_lora_layers(self) -> None: + def convert_to_lora_layers(self, snapshot_base: bool = True) -> None: """ Unified method to convert the transformer to a LoRA transformer. + + snapshot_base=False keeps CPU-backed unmerge snapshots as zero-copy + views of the base weights instead of clones (38 GB of anonymous memory + on H3's DiT). Resident accelerator layers still retain owned CPU + snapshots because they cannot be rebound to the CPU merge cache. """ if self.lora_initialized: return @@ -400,6 +413,7 @@ class LoRAPipeline(ComposedPipelineBase): "transformer", self.lora_layers, check_exclude=True, + snapshot_base=snapshot_base, ) logger.info("Converted %d layers to LoRA layers", converted_count) @@ -412,6 +426,7 @@ class LoRAPipeline(ComposedPipelineBase): self.modules["transformer_2"], "transformer_2", self.lora_layers_transformer_2, + snapshot_base=snapshot_base, check_exclude=True, ) logger.info( @@ -424,6 +439,7 @@ class LoRAPipeline(ComposedPipelineBase): self.modules["fake_score_transformer"], "fake_score_transformer", self.lora_layers_critic, + snapshot_base=snapshot_base, check_exclude=False, ) logger.info( @@ -607,6 +623,7 @@ class LoRAPipeline(ComposedPipelineBase): strengths: list[float], clear_existing: bool = False, merge_weights: bool = True, + merge_cache: LoraMergeCache | None = None, ) -> int: """ Apply LoRA weights to the given lora_layers. Supports multiple LoRA adapters. @@ -667,16 +684,19 @@ class LoRAPipeline(ComposedPipelineBase): layer.lora_rank = inferred_rank layer.lora_alpha = inferred_alpha + use_cache = merge_cache is not None and merge_weights layer.set_lora_weights( self.lora_adapters[nickname][lora_A_name], self.lora_adapters[nickname][lora_B_name], lora_path=path, strength=lora_strength, - merge_weights=merge_weights, + merge_weights=merge_weights and not use_cache, clear_existing=( clear_existing and idx == 0 ), # Only clear on first LoRA ) + if use_cache and idx == len(lora_nicknames) - 1: + self._merge_via_cache(name, layer, merge_cache) adapted_count += 1 applied_count_by_adapter[idx] += 1 else: @@ -912,10 +932,15 @@ class LoRAPipeline(ComposedPipelineBase): merge_weights: bool | None = None, merge_mode: str | None = None, lora_alpha: int | None | list[int | None] = None, + cache_merged: bool = False, ): # type: ignore """ Load LoRA adapter(s) into the pipeline and apply them to the specified transformer(s). Supports both single LoRA (backward compatible) and multiple LoRA adapters. + + cache_merged re-homes merged weights to a file-backed store so they + stop costing anonymous host memory; pass it only for the startup + (static) adapter, where the merged combination is stable. """ merge_mode = self._resolve_lora_merge_mode(merge_weights, merge_mode) @@ -1055,6 +1080,13 @@ class LoRAPipeline(ComposedPipelineBase): tgt_strengths, ) if count is None: + merge_cache = self._merge_cache_for( + module_name, + lora_layers_dict, + tgt_paths, + tgt_strengths, + enabled=cache_merged and effective_merge_weights, + ) count = self._apply_lora_to_layers( lora_layers_dict, tgt_nicknames, @@ -1063,7 +1095,12 @@ class LoRAPipeline(ComposedPipelineBase): tgt_strengths, clear_existing=True, merge_weights=effective_merge_weights, + merge_cache=merge_cache, ) + if merge_cache is not None: + merge_cache.finalize( + {"module": module_name, "paths": tgt_paths} + ) adapted_count += count self.cur_adapter_name[module_name] = merged_name self.cur_adapter_path[module_name] = ",".join( @@ -1092,6 +1129,63 @@ class LoRAPipeline(ComposedPipelineBase): merge_mode, ) + def _merge_via_cache(self, name, layer, merge_cache) -> None: + """Merge one layer through the cache instead of in place. + + The in-place merge copy-on-writes the checkpoint mapping — the whole + component's bytes become anonymous host memory, and so does the + clone() snapshot the layer keeps for unmerging. Going through the + cache leaves the base storage untouched: the layer computes the + merged bytes, the cache holds them file-backed, and the layer adopts + the mapping. If the cache cannot serve or take the bytes, fall back + to the in-place merge — correctness first, memory second. + """ + base_view = layer.weight.data + mapped = merge_cache.get(name, base_view.shape, base_view.dtype) + if mapped is None: + mapped = merge_cache.put(name, layer.compute_merged_weight()) + if mapped is None: + layer.merge_lora_weights() + return + layer.install_merged_weight(mapped, base_view) + + def _merge_cache_for( + self, + module_name: str, + lora_layers: dict[str, BaseLayerWithLoRA], + lora_paths: list[str | None], + strengths: list[float], + enabled: bool, + ) -> LoraMergeCache | None: + if not enabled or envs.SGLANG_DIFFUSION_DISABLE_LORA_MERGE_CACHE: + return None + if any(path is None for path in lora_paths): + return None + if any(layer.weight.device.type != "cpu" for layer in lora_layers.values()): + # Cache entries are CPU mappings. Rebinding a resident accelerator + # parameter to one would leave the module split across devices. + return None + if dist.is_initialized() and dist.get_world_size() > 1: + # Sharded weights would need per-rank stores; not worth it until a + # multi-GPU consumer deployment exists. + return None + adapters = [ + (path, strength, self.server_args.lora_alpha) + for path, strength in zip(lora_paths, strengths) + ] + key = lora_merge_cache_key([self.server_args.model_path, module_name], adapters) + expected = sum( + layer.weight.numel() * layer.weight.element_size() + for layer in lora_layers.values() + ) + store = LoraMergeCache(key, expected) + if store.is_complete(): + logger.info( + "LoRA merge cache found for %s; adopting instead of merging", + module_name, + ) + return store + def deactivate_lora_weights(self, target: str = "all") -> None: """ Disable LoRA for the specified target, regardless of whether weights were diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_merge_cache.py b/python/sglang/multimodal_gen/test/unit/test_lora_merge_cache.py new file mode 100644 index 000000000..f36b5f4a1 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_lora_merge_cache.py @@ -0,0 +1,77 @@ +"""The merge cache is a pure byte vault: tensors in, mapped tensors out. + +What matters: put round-trips the exact merged bytes and returns a mapping, +a complete cache serves the same bytes back, a mismatched entry is refused, +disk shortage returns None instead of raising, and the key separates +combinations that must not share files. +""" + +import pytest +import torch + +from sglang.multimodal_gen.runtime.pipelines_core.lora.lora_merge_cache import ( + LoraMergeCache, + lora_merge_cache_key, +) + + +@pytest.fixture(autouse=True) +def _cache_root(monkeypatch, tmp_path): + monkeypatch.setenv("SGLANG_DIFFUSION_CACHE_ROOT", str(tmp_path / "cache")) + + +def test_put_round_trips_and_returns_a_mapping(): + merged = torch.randn(16, 16) + cache = LoraMergeCache("k1", expected_bytes=merged.numel() * 4) + + mapped = cache.put("blocks.0.linear", merged) + assert mapped is not None + assert torch.equal(mapped, merged) + cache.finalize() + + second = LoraMergeCache("k1", expected_bytes=0) + assert second.is_complete() + served = second.get("blocks.0.linear", merged.shape, merged.dtype) + assert served is not None + assert torch.equal(served, merged) + + +def test_an_incomplete_cache_is_not_complete(): + cache = LoraMergeCache("k2", expected_bytes=64) + assert cache.put("a", torch.randn(4, 4)) is not None + # no finalize -> no manifest + assert not LoraMergeCache("k2", expected_bytes=0).is_complete() + + +def test_a_mismatched_entry_is_refused(): + cache = LoraMergeCache("k3", expected_bytes=64) + assert cache.put("a", torch.randn(4, 4)) is not None + cache.finalize() + + second = LoraMergeCache("k3", expected_bytes=0) + assert second.is_complete() + assert second.get("a", torch.Size([8, 8]), torch.float32) is None + + +def test_disk_shortage_returns_none(monkeypatch): + import shutil as _shutil + from types import SimpleNamespace + + monkeypatch.setattr( + _shutil, "disk_usage", lambda _: SimpleNamespace(free=1, total=1, used=0) + ) + cache = LoraMergeCache("k4", expected_bytes=1 << 40) + assert cache.put("a", torch.randn(4, 4)) is None + + +def test_the_key_separates_combinations(tmp_path): + lora = tmp_path / "adapter.safetensors" + lora.write_bytes(b"x" * 128) + base = ["/models/h3", "transformer"] + k = lora_merge_cache_key(base, [(str(lora), 1.0, None)]) + assert k != lora_merge_cache_key(base, [(str(lora), 0.5, None)]) + assert k != lora_merge_cache_key(base, [(str(lora), 1.0, 32)]) + assert k != lora_merge_cache_key( + ["/models/h3", "transformer_2"], [(str(lora), 1.0, None)] + ) + assert k == lora_merge_cache_key(base, [(str(lora), 1.0, None)]) diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py b/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py index c15db3e1b..28b2a2340 100644 --- a/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py +++ b/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py @@ -5,7 +5,11 @@ from unittest.mock import patch import torch -from sglang.multimodal_gen.runtime.layers.lora.linear import BaseLayerWithLoRA +from sglang.multimodal_gen.runtime.layers.lora.linear import ( + BaseLayerWithLoRA, + _use_owned_base_snapshot, + wrap_with_lora_layer, +) from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_lora @@ -24,7 +28,11 @@ def _make_layer() -> BaseLayerWithLoRA: def _make_pipeline(layer: BaseLayerWithLoRA) -> _TestLoRAPipeline: pipeline = object.__new__(_TestLoRAPipeline) pipeline.modules = {"transformer": torch.nn.Module()} - pipeline.server_args = SimpleNamespace(lora_merge_mode="dynamic") + pipeline.server_args = SimpleNamespace( + lora_alpha=None, + lora_merge_mode="dynamic", + model_path="/model", + ) pipeline.lora_initialized = True pipeline.lora_adapters = defaultdict(dict) pipeline.loaded_adapter_paths = {"adapter": "/adapter"} @@ -43,6 +51,49 @@ def _make_pipeline(layer: BaseLayerWithLoRA) -> _TestLoRAPipeline: return pipeline +def test_merge_cache_only_accepts_cpu_backed_weights(): + pipeline = _make_pipeline(_make_layer()) + cpu_cache = pipeline._merge_cache_for( + "transformer", + pipeline.lora_layers, + ["/adapter"], + [1.0], + enabled=True, + ) + assert cpu_cache is not None + + resident_layer = BaseLayerWithLoRA( + torch.nn.Linear(2, 2, bias=False, device="meta"), snapshot_base=False + ) + resident_cache = pipeline._merge_cache_for( + "transformer", + {"linear": resident_layer}, + ["/adapter"], + [1.0], + enabled=True, + ) + assert resident_cache is None + + +def test_zero_copy_snapshot_is_limited_to_cpu_backed_layers(): + assert not _use_owned_base_snapshot(False, "cpu") + assert not _use_owned_base_snapshot(False, "meta") + assert _use_owned_base_snapshot(False, "cuda") + assert _use_owned_base_snapshot(True, "cpu") + + cpu_layer = wrap_with_lora_layer( + torch.nn.Linear(2, 2, bias=False), snapshot_base=False + ) + assert cpu_layer is not None + assert cpu_layer._base_is_view + + meta_layer = wrap_with_lora_layer( + torch.nn.Linear(2, 2, bias=False, device="meta"), snapshot_base=False + ) + assert meta_layer is not None + assert meta_layer._base_is_view + + def test_dynamic_lora_reactivates_cached_layers_without_weight_update_context(): layer = _make_layer() pipeline = _make_pipeline(layer)