[lora] Fix overlap loading for cancelled requests (#25413)
This commit is contained in:
@@ -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,19 +55,26 @@ 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
|
||||||
|
|
||||||
|
# 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.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)
|
torch.cuda.current_stream().wait_event(event)
|
||||||
del self.lora_to_overlap_load_event[lora_id]
|
del self.lora_to_overlap_load_event[lora_id]
|
||||||
|
|
||||||
return LoRAOverlapLoadStatus.LOADED
|
|
||||||
|
|
||||||
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]]
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user