[Fix] Clear full-to-SWA mapping with index_fill_ to avoid a blocking H2D copy (#35773)

This commit is contained in:
Liangsheng Yin
2026-08-21 01:23:26 -07:00
committed by GitHub
parent 4c98759c73
commit 896acc8860
6 changed files with 53 additions and 9 deletions
@@ -226,6 +226,43 @@ class TestSWA(unittest.TestCase):
allocator.free_swa(full_indices[1:2])
self.assertEqual(allocator.swa_available_size(), 16)
@unittest.skipUnless(torch.cuda.is_available(), "sync detection needs CUDA")
def test_clearing_the_mapping_does_not_synchronize(self):
"""Clearing the full-to-SWA mapping must not block the stream; writing a
host-resident scalar into it does.
"""
_, allocator, _ = _build_swa_tree(is_eagle=False)
full_indices = _swa_alloc(allocator, 4)
mapping = allocator.full_to_swa_index_mapping
# Warm up outside the window: a first-time cudaMalloc can synchronize on
# its own, which the detector would report as this call's fault.
allocator.clear_full_to_swa_mapping(full_indices)
def sync_error(fn):
torch.cuda.synchronize()
torch.cuda.set_sync_debug_mode("error")
try:
fn()
except RuntimeError as exc:
return exc
finally:
torch.cuda.set_sync_debug_mode("default")
torch.cuda.synchronize()
return None
# Gate on the pre-fix form: a detector blind to this sync class would pass
# the assert below no matter how the mapping is cleared.
pre_fix_error = sync_error(
lambda: mapping.__setitem__(full_indices.to(torch.int64), 0)
)
if pre_fix_error is None:
self.skipTest("sync debug mode does not flag a blocking H2D copy here")
self.assertIsNone(
sync_error(lambda: allocator.clear_full_to_swa_mapping(full_indices))
)
def test_free_swa_group_owns_deferred_indices(self):
_, allocator, _ = _build_swa_tree(
is_eagle=False,
@@ -6147,9 +6147,7 @@ class TestUnifiedRadixCacheActionRouting(CustomTestCase):
alloc.translate_loc_from_full_to_swa.assert_called_once_with(incoming_full)
alloc.set_full_to_swa_mapping.assert_called_once_with(kept_full, swa_value)
# the incoming full's stale mapping is cleared, then its slot freed (full-only)
key, val = alloc.full_to_swa_index_mapping.__setitem__.call_args.args
self.assertTrue(torch.equal(key, incoming_full))
self.assertEqual(val, 0)
alloc.clear_full_to_swa_mapping.assert_called_once_with(incoming_full)
alloc.full_attn_allocator.free.assert_called_once_with(incoming_full)
alloc.free.assert_not_called()
cache.tree_core.set_component_device_value.assert_called_once_with(