[lora] Fix overlap loading for cancelled requests (#25413)

This commit is contained in:
Erik Wijmans
2026-05-25 15:18:36 +09:00
committed by GitHub
parent 2bd3ac0b5d
commit 87e69d57c4
2 changed files with 98 additions and 9 deletions
+20 -9
View File
@@ -35,6 +35,10 @@ class LoRAOverlapLoader:
Check a LoRA adapter's asynchronous load status, and try to load it if there's capacity Check a LoRA adapter's asynchronous load status, and try to load it if there's capacity
in the memory pool. Returns whether or not the adapter has been loaded. in the memory pool. Returns whether or not the adapter has been loaded.
""" """
# Drain completed async loads before status/capacity checks so finished
# adapters no longer count as in-flight.
self._drain_completed_overlap_loads()
lora_pipeline_load_status = self._check_overlap_load_status(lora_id) lora_pipeline_load_status = self._check_overlap_load_status(lora_id)
if lora_pipeline_load_status == LoRAOverlapLoadStatus.LOADING: if lora_pipeline_load_status == LoRAOverlapLoadStatus.LOADING:
return False return False
@@ -51,18 +55,25 @@ class LoRAOverlapLoader:
def _check_overlap_load_status( def _check_overlap_load_status(
self, lora_id: Optional[str] self, lora_id: Optional[str]
) -> LoRAOverlapLoadStatus: ) -> LoRAOverlapLoadStatus:
if lora_id not in self.lora_to_overlap_load_event: if lora_id in self.lora_to_overlap_load_event:
return LoRAOverlapLoadStatus.NOT_LOADED
event = self.lora_to_overlap_load_event[lora_id]
if not event.query():
return LoRAOverlapLoadStatus.LOADING return LoRAOverlapLoadStatus.LOADING
torch.cuda.current_stream().wait_event(event) # After completed events have been drained, a memory-pool entry with no
del self.lora_to_overlap_load_event[lora_id] # pending event is safe to use on the current stream.
if lora_id in self.lora_manager.memory_pool.uid_to_buffer_id:
return LoRAOverlapLoadStatus.LOADED
return LoRAOverlapLoadStatus.LOADED return LoRAOverlapLoadStatus.NOT_LOADED
def _drain_completed_overlap_loads(self) -> None:
completed_loads = [
(lora_id, event)
for lora_id, event in self.lora_to_overlap_load_event.items()
if event.query()
]
for lora_id, event in completed_loads:
torch.cuda.current_stream().wait_event(event)
del self.lora_to_overlap_load_event[lora_id]
def _try_start_overlap_load( def _try_start_overlap_load(
self, lora_id: Optional[str], running_loras: set[Optional[str]] self, lora_id: Optional[str], running_loras: set[Optional[str]]
@@ -65,7 +65,10 @@ class TestLoRAOverlapLoaderUnitTests(CustomTestCase):
self.mock_lora_manager = MagicMock(spec=LoRAManager) self.mock_lora_manager = MagicMock(spec=LoRAManager)
self.mock_lora_manager.device = "cuda:0" self.mock_lora_manager.device = "cuda:0"
self.mock_lora_manager.memory_pool = MagicMock()
self.mock_lora_manager.memory_pool.uid_to_buffer_id = {}
self.mock_lora_manager.validate_lora_batch.return_value = True self.mock_lora_manager.validate_lora_batch.return_value = True
self.mock_lora_manager.fetch_new_loras.side_effect = self._mark_loras_loaded
def tearDown(self): def tearDown(self):
self.torch_patcher.stop() self.torch_patcher.stop()
@@ -73,11 +76,85 @@ class TestLoRAOverlapLoaderUnitTests(CustomTestCase):
def _create_loader(self) -> LoRAOverlapLoader: def _create_loader(self) -> LoRAOverlapLoader:
return LoRAOverlapLoader(cast(LoRAManager, self.mock_lora_manager)) return LoRAOverlapLoader(cast(LoRAManager, self.mock_lora_manager))
def _mark_loras_loaded(self, new_loras, _loras_to_be_loaded):
for lora_id in new_loras:
self.mock_lora_manager.memory_pool.uid_to_buffer_id[lora_id] = len(
self.mock_lora_manager.memory_pool.uid_to_buffer_id
)
def _create_mock_event(self, query_return: bool = False) -> MagicMock: def _create_mock_event(self, query_return: bool = False) -> MagicMock:
event = MagicMock(spec=CudaEvent) event = MagicMock(spec=CudaEvent)
event.query.return_value = query_return event.query.return_value = query_return
return event return event
def test_completed_stale_loads_are_reaped_before_capacity_check(self):
loader = self._create_loader()
events = [
self._create_mock_event(query_return=True),
self._create_mock_event(query_return=False),
]
self.mock_device_module.Event.side_effect = events
self.mock_lora_manager.validate_lora_batch.side_effect = (
lambda lora_ids: len(lora_ids) <= 1
)
self.assertTrue(
loader._try_start_overlap_load("stale_lora", running_loras=set())
)
self.assertIn("stale_lora", loader.lora_to_overlap_load_event)
self.mock_lora_manager.fetch_new_loras.reset_mock()
result = loader.try_overlap_load_lora("new_lora", running_loras=set())
self.assertFalse(result)
self.assertNotIn("stale_lora", loader.lora_to_overlap_load_event)
self.assertIn("new_lora", loader.lora_to_overlap_load_event)
self.mock_lora_manager.fetch_new_loras.assert_called_once_with(
{"new_lora"}, set()
)
def test_loaded_lora_reused_after_stale_event_drain(self):
loader = self._create_loader()
self.mock_lora_manager.memory_pool = MagicMock()
self.mock_lora_manager.memory_pool.uid_to_buffer_id = {}
events = [
self._create_mock_event(query_return=True),
self._create_mock_event(query_return=False),
]
self.mock_device_module.Event.side_effect = events
self.mock_lora_manager.validate_lora_batch.side_effect = (
lambda lora_ids: len(lora_ids) <= 2
)
self.assertTrue(loader._try_start_overlap_load("lora_A", running_loras=set()))
self.assertIn("lora_A", loader.lora_to_overlap_load_event)
self.mock_lora_manager.fetch_new_loras.reset_mock()
self.assertFalse(loader.try_overlap_load_lora("lora_B", running_loras=set()))
self.assertNotIn("lora_A", loader.lora_to_overlap_load_event)
self.assertIn("lora_B", loader.lora_to_overlap_load_event)
self.mock_lora_manager.fetch_new_loras.assert_called_once_with(
{"lora_B"}, set()
)
self.mock_lora_manager.fetch_new_loras.reset_mock()
self.assertTrue(loader.try_overlap_load_lora("lora_A", running_loras=set()))
self.assertIn("lora_B", loader.lora_to_overlap_load_event)
self.mock_lora_manager.fetch_new_loras.assert_not_called()
def test_pending_lora_load_must_complete_even_if_memory_pool_has_slot(self):
loader = self._create_loader()
self.mock_lora_manager.memory_pool = MagicMock()
self.mock_lora_manager.memory_pool.uid_to_buffer_id = {"lora_A": 0}
loader.lora_to_overlap_load_event["lora_A"] = self._create_mock_event(False)
result = loader.try_overlap_load_lora("lora_A", running_loras=set())
self.assertFalse(result)
self.mock_lora_manager.fetch_new_loras.assert_not_called()
self.assertIn("lora_A", loader.lora_to_overlap_load_event)
def test_full_lifecycle_single_lora_load(self): def test_full_lifecycle_single_lora_load(self):
loader = self._create_loader() loader = self._create_loader()
@@ -131,6 +208,7 @@ class TestLoRAOverlapLoaderUnitTests(CustomTestCase):
# First lora completes, freeing capacity # First lora completes, freeing capacity
loader.lora_to_overlap_load_event["lora_0"].query.return_value = True loader.lora_to_overlap_load_event["lora_0"].query.return_value = True
loader._drain_completed_overlap_loads()
self.assertEqual( self.assertEqual(
loader._check_overlap_load_status("lora_0"), LoRAOverlapLoadStatus.LOADED loader._check_overlap_load_status("lora_0"), LoRAOverlapLoadStatus.LOADED
) )