From 62ba9648482e1b0a253187c8dc6ac5d23bbbd403 Mon Sep 17 00:00:00 2001 From: jianzhao-xu <978716854@qq.com> Date: Mon, 21 Sep 2026 11:19:07 +0800 Subject: [PATCH] Fix: post-load staging regression breaks offload meta/sharded_gpu modes (#38779) --- python/sglang/srt/model_loader/post_load.py | 12 +++++++--- .../unit/test_quantization_post_load.py | 23 +++++++++++++++---- 2 files changed, 27 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/model_loader/post_load.py b/python/sglang/srt/model_loader/post_load.py index 6d0110f08..91c5ba8bb 100644 --- a/python/sglang/srt/model_loader/post_load.py +++ b/python/sglang/srt/model_loader/post_load.py @@ -134,13 +134,16 @@ def stage_module_for_post_load( owner_origins: dict[int, set[torch.device]] = {} module_origins: set[torch.device] = set() tensor_states: dict[int, _TensorState] = {} + # Offloader-parked parameters (see sglang.srt.utils.offloader) live on the + # meta device through post-load processing; kernels are expected to skip + # them. Pass them through untouched, matching the pre-staging behaviour. + meta_tensor_ids: set[int] = set() # snapshot and validate all state before moving any tensor for owner, registry_name, name, tensor in _iter_registered_tensors(module): if tensor.is_meta: - raise RuntimeError( - f"Cannot post-process meta tensor {type(owner).__name__}.{name}" - ) + meta_tensor_ids.add(id(tensor)) + continue state = tensor_states.get(id(tensor)) if state is None: state = _TensorState(tensor, tensor.data, tensor.device) @@ -167,6 +170,9 @@ def stage_module_for_post_load( next(iter(module_origins)) if len(module_origins) == 1 else None ) for owner, registry_name, name, tensor in _iter_registered_tensors(module): + if id(tensor) in meta_tensor_ids: + # Parked by the offloader before staging; leave it as-is. + continue key = _slot_key(owner, registry_name, name) original_state = original_slots.get(key) if original_state is None: diff --git a/test/registered/unit/test_quantization_post_load.py b/test/registered/unit/test_quantization_post_load.py index 11a7fc1c5..11cf2e6bd 100644 --- a/test/registered/unit/test_quantization_post_load.py +++ b/test/registered/unit/test_quantization_post_load.py @@ -140,18 +140,31 @@ class TestModulePostLoadValidation(unittest.TestCase): self.assertEqual(module.weight.device, PROCESS_DEVICE) - def test_rejects_meta_before_moving_other_state(self): + def test_parks_existing_meta_tensor_and_stages_the_rest(self): module = nn.Module() module.meta_weight = nn.Parameter(torch.empty(1, device="meta")) module.register_buffer("scale", torch.ones(1)) scale_ptr = module.scale.data_ptr() + meta_weight = module.meta_weight + + with stage_module_for_post_load(module, torch.device("cpu")): + # offloader-parked meta state passes through untouched + self.assertTrue(module.meta_weight.is_meta) + module.scale.add_(1) + + self.assertIs(module.meta_weight, meta_weight) + self.assertTrue(module.meta_weight.is_meta) + self.assertEqual(module.scale.device.type, "cpu") + self.assertEqual(module.scale.data_ptr(), scale_ptr) + torch.testing.assert_close(module.scale, torch.full((1,), 2.0)) + + def test_rejects_meta_produced_by_post_load_processing(self): + module = nn.Module() + module.register_buffer("scale", torch.ones(1)) with self.assertRaisesRegex(RuntimeError, "meta tensor"): with stage_module_for_post_load(module, torch.device("cpu")): - self.fail("context should not be entered") - - self.assertEqual(module.scale.device.type, "cpu") - self.assertEqual(module.scale.data_ptr(), scale_ptr) + module.new_weight = nn.Parameter(torch.empty(1, device="meta")) if __name__ == "__main__":