diff --git a/python/sglang/srt/lora/lora_overlap_loader.py b/python/sglang/srt/lora/lora_overlap_loader.py index bc7b3dd71..6d5845ba0 100644 --- a/python/sglang/srt/lora/lora_overlap_loader.py +++ b/python/sglang/srt/lora/lora_overlap_loader.py @@ -35,6 +35,10 @@ class LoRAOverlapLoader: 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. """ + # 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) if lora_pipeline_load_status == LoRAOverlapLoadStatus.LOADING: return False @@ -51,18 +55,25 @@ class LoRAOverlapLoader: def _check_overlap_load_status( self, lora_id: Optional[str] ) -> LoRAOverlapLoadStatus: - if lora_id not in self.lora_to_overlap_load_event: - return LoRAOverlapLoadStatus.NOT_LOADED - - event = self.lora_to_overlap_load_event[lora_id] - - if not event.query(): + if lora_id in self.lora_to_overlap_load_event: return LoRAOverlapLoadStatus.LOADING - torch.cuda.current_stream().wait_event(event) - del self.lora_to_overlap_load_event[lora_id] + # After completed events have been drained, a memory-pool entry with no + # 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( self, lora_id: Optional[str], running_loras: set[Optional[str]] diff --git a/test/registered/lora/test_lora_overlap_loading.py b/test/registered/lora/test_lora_overlap_loading.py index ea7765cf6..2d187b20e 100644 --- a/test/registered/lora/test_lora_overlap_loading.py +++ b/test/registered/lora/test_lora_overlap_loading.py @@ -65,7 +65,10 @@ class TestLoRAOverlapLoaderUnitTests(CustomTestCase): self.mock_lora_manager = MagicMock(spec=LoRAManager) 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.fetch_new_loras.side_effect = self._mark_loras_loaded def tearDown(self): self.torch_patcher.stop() @@ -73,11 +76,85 @@ class TestLoRAOverlapLoaderUnitTests(CustomTestCase): def _create_loader(self) -> LoRAOverlapLoader: 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: event = MagicMock(spec=CudaEvent) event.query.return_value = query_return 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): loader = self._create_loader() @@ -131,6 +208,7 @@ class TestLoRAOverlapLoaderUnitTests(CustomTestCase): # First lora completes, freeing capacity loader.lora_to_overlap_load_event["lora_0"].query.return_value = True + loader._drain_completed_overlap_loads() self.assertEqual( loader._check_overlap_load_status("lora_0"), LoRAOverlapLoadStatus.LOADED )