Fix: post-load staging regression breaks offload meta/sharded_gpu modes (#38779)
This commit is contained in:
@@ -134,13 +134,16 @@ def stage_module_for_post_load(
|
|||||||
owner_origins: dict[int, set[torch.device]] = {}
|
owner_origins: dict[int, set[torch.device]] = {}
|
||||||
module_origins: set[torch.device] = set()
|
module_origins: set[torch.device] = set()
|
||||||
tensor_states: dict[int, _TensorState] = {}
|
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
|
# snapshot and validate all state before moving any tensor
|
||||||
for owner, registry_name, name, tensor in _iter_registered_tensors(module):
|
for owner, registry_name, name, tensor in _iter_registered_tensors(module):
|
||||||
if tensor.is_meta:
|
if tensor.is_meta:
|
||||||
raise RuntimeError(
|
meta_tensor_ids.add(id(tensor))
|
||||||
f"Cannot post-process meta tensor {type(owner).__name__}.{name}"
|
continue
|
||||||
)
|
|
||||||
state = tensor_states.get(id(tensor))
|
state = tensor_states.get(id(tensor))
|
||||||
if state is None:
|
if state is None:
|
||||||
state = _TensorState(tensor, tensor.data, tensor.device)
|
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
|
next(iter(module_origins)) if len(module_origins) == 1 else None
|
||||||
)
|
)
|
||||||
for owner, registry_name, name, tensor in _iter_registered_tensors(module):
|
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)
|
key = _slot_key(owner, registry_name, name)
|
||||||
original_state = original_slots.get(key)
|
original_state = original_slots.get(key)
|
||||||
if original_state is None:
|
if original_state is None:
|
||||||
|
|||||||
@@ -140,18 +140,31 @@ class TestModulePostLoadValidation(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertEqual(module.weight.device, PROCESS_DEVICE)
|
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 = nn.Module()
|
||||||
module.meta_weight = nn.Parameter(torch.empty(1, device="meta"))
|
module.meta_weight = nn.Parameter(torch.empty(1, device="meta"))
|
||||||
module.register_buffer("scale", torch.ones(1))
|
module.register_buffer("scale", torch.ones(1))
|
||||||
scale_ptr = module.scale.data_ptr()
|
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 self.assertRaisesRegex(RuntimeError, "meta tensor"):
|
||||||
with stage_module_for_post_load(module, torch.device("cpu")):
|
with stage_module_for_post_load(module, torch.device("cpu")):
|
||||||
self.fail("context should not be entered")
|
module.new_weight = nn.Parameter(torch.empty(1, device="meta"))
|
||||||
|
|
||||||
self.assertEqual(module.scale.device.type, "cpu")
|
|
||||||
self.assertEqual(module.scale.data_ptr(), scale_ptr)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user