From f294d51a718d8621abdb67951dcd40d23c371d94 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 24 Aug 2026 11:11:35 +0800 Subject: [PATCH] [diffusion] fix: fix a refit key error on mapped weights, and stop claiming strides the reload discards (#35832) Co-authored-by: Claude Opus 5 --- .../memory_managers/layerwise_offload.py | 22 +++++- .../test/unit/test_layerwise_offload.py | 76 +++++++++++++++++-- 2 files changed, 89 insertions(+), 9 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py index 94827f4fa..fb1e37831 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py @@ -711,7 +711,14 @@ class LayerwiseOffloadManager: contiguous_weights: List[Tuple[str, torch.Tensor, torch.Tensor]] = [] for name, weight in weights: local_weight = self._to_local_tensor(weight) - if hosting == "mapped" and self._mapped_regions.holds(local_weight): + if ( + hosting == "mapped" + and local_weight.is_contiguous() + and self._mapped_regions.holds(local_weight) + ): + # Only a contiguous view can stay mapped: the reload + # path allocates contiguous and would drop any other + # layout without saying so. # Already a view into the checkpoint. Copying it would # add a second copy of bytes the page cache holds # anyway, and that copy is what does not fit. @@ -727,7 +734,6 @@ class LayerwiseOffloadManager: self._weight_metadata[layer_idx][name] = { "dtype": local_weight.dtype, "shape": tuple(local_weight.shape), - "stride": local_weight.stride(), "preserve_strides": False, "mapped": True, } @@ -1298,7 +1304,17 @@ class LayerwiseOffloadManager: ) dtype = meta["dtype"] - if meta.get("preserve_strides", False): + if meta.get("mapped", False): + # The mapping is a read-only view of the checkpoint, so the new + # values cannot be written into it. Own the storage from here + # on; every reader of this store copies out of whatever tensor + # it holds. This trades mapped bytes for anonymous ones on the + # configuration that chose mapping because host memory was + # short, so it costs the updated weight's bytes. + self._mapped_cpu_weights[layer_idx][name] = ( + local_loaded_weight.detach().to(dtype=dtype).contiguous() + ) + elif meta.get("preserve_strides", False): self._strided_cpu_weights[layer_idx][name].copy_( local_loaded_weight.to(dtype=dtype) ) 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 01ecb819b..1fbc65cbc 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -1323,14 +1323,26 @@ class _FileBackedBlock(torch.nn.Module): self.weight = torch.nn.Parameter(mapped.reshape(8, 8), requires_grad=False) +class _TransposedFileBackedBlock(torch.nn.Module): + """A mapped weight whose layout is not contiguous, as an FP8 weight is.""" + + def __init__(self, path: pathlib.Path) -> None: + super().__init__() + path.write_bytes(b"\x00" * (64 * 4)) + mapped = torch.from_file(str(path), shared=True, size=64, dtype=torch.float32) + self.weight = torch.nn.Parameter(mapped.reshape(8, 8).t(), requires_grad=False) + + class _FileBackedModel(torch.nn.Module): - def __init__(self, path: pathlib.Path, num_blocks: int = 1) -> None: + def __init__( + self, + path: pathlib.Path, + num_blocks: int = 1, + block_cls=_FileBackedBlock, + ) -> None: super().__init__() self.blocks = torch.nn.ModuleList( - [ - _FileBackedBlock(path.with_name(f"{path.name}.{i}")) - for i in range(num_blocks) - ] + [block_cls(path.with_name(f"{path.name}.{i}")) for i in range(num_blocks)] ) @@ -1450,6 +1462,7 @@ def _mapped_manager( available_bytes=None, num_blocks=1, pin_budget_bytes=None, + block_cls=_FileBackedBlock, ): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule @@ -1460,7 +1473,9 @@ def _mapped_manager( monkeypatch.setattr( host_memory_budget, "host_memory_available_bytes", lambda: available_bytes ) - model = _FileBackedModel(tmp_path / "weights.bin", num_blocks=num_blocks) + model = _FileBackedModel( + tmp_path / "weights.bin", num_blocks=num_blocks, block_cls=block_cls + ) return LayerwiseOffloadManager( model=model, layers_attr_str="blocks", @@ -1671,6 +1686,55 @@ def test_mapped_weights_are_visible_to_checksums(tmp_path, monkeypatch): assert "blocks.0.weight" in names +def test_a_non_contiguous_mapped_weight_keeps_its_layout(tmp_path, monkeypatch): + """Staying mapped costs the layout, so a strided weight must not stay. + + The reload path allocates with `torch.empty(shape)` and copies, which is + layout-agnostic: values survive, strides do not. ModelOpt FP8 calls its + transposed layout a correctness requirement, so such a weight has to take + the strided path even when its storage is a mapping the copies cannot + afford. + """ + if not pathlib.Path("/proc/self/maps").exists(): + pytest.skip("needs /proc to tell a mapping from anonymous memory") + manager = _mapped_manager( + tmp_path, + monkeypatch, + available_gib=0.001, + pin_budget_bytes=0, + block_cls=_TransposedFileBackedBlock, + ) + name = "blocks.0.weight" + assert name not in manager._mapped_cpu_weights[0], ( + "a non-contiguous weight stayed mapped, so its layout is dropped " + "on reload without any error" + ) + stored = manager._strided_cpu_weights[0][name] + assert not stored.is_contiguous() + assert stored.stride() == (1, 8) + assert manager._weight_metadata[0][name]["preserve_strides"] is True + + +def test_refitting_a_mapped_weight_updates_the_store(tmp_path, monkeypatch): + """A refit must reach a mapped weight without writing to the checkpoint.""" + if not pathlib.Path("/proc/self/maps").exists(): + pytest.skip("needs /proc to tell a mapping from anonymous memory") + path = tmp_path / "weights.bin.0" + manager = _mapped_manager( + tmp_path, monkeypatch, available_gib=0.001, pin_budget_bytes=0 + ) + name = "blocks.0.weight" + assert manager._weight_metadata[0][name]["mapped"] is True + on_disk_before = path.read_bytes() + + new_weight = torch.full((8, 8), 3.0) + updated = manager.update_cpu_weights({name: new_weight}) + + assert updated == {name} + assert torch.equal(manager._mapped_cpu_weights[0][name], new_weight) + assert path.read_bytes() == on_disk_before, "the checkpoint was written to" + + def test_layerwise_tuning_defaults_match_the_group(): """No per-component entry: the DiT group keeps its knobs, auxiliaries do not.""" args = _server_args(