[TestFix] use unit tests for LoRA overlap loading tests (#18140)

This commit is contained in:
Glen Liu
2026-02-03 01:06:50 -05:00
committed by GitHub
parent f032c4f3d6
commit fe57a887b1

View File

@@ -11,23 +11,26 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""
End-to-end tests for the --enable-lora-overlap-loading server argument.
"""
import multiprocessing as mp
import unittest
from typing import cast
from unittest.mock import MagicMock, patch
from torch.cuda import Event as CudaEvent
from torch.cuda import Stream as CudaStream
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.srt.lora.lora_overlap_loader import LoRAOverlapLoader, LoRAOverlapLoadStatus
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.lora_utils import (
CI_MULTI_LORA_MODELS,
run_lora_batch_splitting_equivalence_test,
run_lora_multiple_batch_on_model_cases,
)
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=100, suite="stage-b-test-large-1-gpu")
register_amd_ci(est_time=100, suite="stage-b-test-small-1-gpu-amd")
register_cuda_ci(est_time=75, suite="stage-b-test-large-1-gpu")
register_amd_ci(est_time=75, suite="stage-b-test-small-1-gpu-amd")
class TestLoRAOverlapLoading(CustomTestCase):
@@ -36,11 +39,126 @@ class TestLoRAOverlapLoading(CustomTestCase):
CI_MULTI_LORA_MODELS, enable_lora_overlap_loading=True
)
def test_ci_lora_models_multi_batch(self):
run_lora_multiple_batch_on_model_cases(
CI_MULTI_LORA_MODELS, enable_lora_overlap_loading=True
class TestLoRAOverlapLoaderUnitTests(CustomTestCase):
mock_lora_manager: MagicMock
mock_stream: MagicMock
mock_stream_context: MagicMock
mock_device_module: MagicMock
mock_torch: MagicMock
def setUp(self):
self.torch_patcher = patch("sglang.srt.lora.lora_overlap_loader.torch")
self.mock_torch = self.torch_patcher.start()
self.mock_device_module = MagicMock()
self.mock_stream = MagicMock(spec=CudaStream)
self.mock_stream_context = MagicMock()
self.mock_event = MagicMock(spec=CudaEvent)
self.mock_device_module.Stream.return_value = self.mock_stream
self.mock_device_module.stream.return_value = self.mock_stream_context
self.mock_device_module.Event.return_value = self.mock_event
self.mock_torch.get_device_module.return_value = self.mock_device_module
self.mock_torch.cuda.current_stream.return_value = MagicMock(spec=CudaStream)
self.mock_lora_manager = MagicMock(spec=LoRAManager)
self.mock_lora_manager.device = "cuda:0"
self.mock_lora_manager.validate_lora_batch.return_value = True
def tearDown(self):
self.torch_patcher.stop()
def _create_loader(self) -> LoRAOverlapLoader:
return LoRAOverlapLoader(cast(LoRAManager, self.mock_lora_manager))
def _create_mock_event(self, query_return: bool = False) -> MagicMock:
event = MagicMock(spec=CudaEvent)
event.query.return_value = query_return
return event
def test_full_lifecycle_single_lora_load(self):
loader = self._create_loader()
# Initially not loaded
status = loader._check_overlap_load_status("lora_A")
self.assertEqual(status, LoRAOverlapLoadStatus.NOT_LOADED)
# First call starts async load, returns False
result = loader.try_overlap_load_lora("lora_A", running_loras=set())
self.assertFalse(result)
self.assertIn("lora_A", loader.lora_to_overlap_load_event)
self.mock_lora_manager.fetch_new_loras.assert_called_once_with(
{"lora_A"}, set()
)
# Simulate load still in progress - returns False, event persists
loader.lora_to_overlap_load_event["lora_A"].query.return_value = False
result = loader.try_overlap_load_lora("lora_A", running_loras=set())
self.assertFalse(result)
self.assertEqual(
loader._check_overlap_load_status("lora_A"), LoRAOverlapLoadStatus.LOADING
)
# Simulate load complete - returns True, event removed
loader.lora_to_overlap_load_event["lora_A"].query.return_value = True
result = loader.try_overlap_load_lora("lora_A", running_loras=set())
self.assertTrue(result)
self.assertNotIn("lora_A", loader.lora_to_overlap_load_event)
def test_capacity_constraints_block_new_loads(self):
loader = self._create_loader()
events = [self._create_mock_event() for _ in range(4)]
self.mock_device_module.Event.side_effect = events
# Load 3 loras successfully
for i in range(3):
self.assertTrue(
loader._try_start_overlap_load(f"lora_{i}", running_loras=set())
)
self.assertEqual(len(loader.lora_to_overlap_load_event), 3)
# Capacity full - new load blocked
self.mock_lora_manager.validate_lora_batch.return_value = False
self.mock_lora_manager.fetch_new_loras.reset_mock()
result = loader.try_overlap_load_lora("lora_3", running_loras=set())
self.assertFalse(result)
self.mock_lora_manager.fetch_new_loras.assert_not_called()
self.assertNotIn("lora_3", loader.lora_to_overlap_load_event)
# First lora completes, freeing capacity
loader.lora_to_overlap_load_event["lora_0"].query.return_value = True
self.assertEqual(
loader._check_overlap_load_status("lora_0"), LoRAOverlapLoadStatus.LOADED
)
# Now new load succeeds
self.mock_lora_manager.validate_lora_batch.return_value = True
self.assertTrue(loader._try_start_overlap_load("lora_3", running_loras=set()))
def test_validation_includes_pending_and_running_loras(self):
loader = self._create_loader()
events = [self._create_mock_event() for _ in range(5)]
self.mock_device_module.Event.side_effect = events
# Start pending loads
loader._try_start_overlap_load("pending_1", running_loras=set())
loader._try_start_overlap_load("pending_2", running_loras=set())
# Load new lora with running_loras
self.mock_lora_manager.validate_lora_batch.reset_mock()
running = {"running_1", "running_2"}
loader.try_overlap_load_lora("new_lora", running_loras=running)
# Validation should include: pending + running + new
call_args = self.mock_lora_manager.validate_lora_batch.call_args[0][0]
expected = {"pending_1", "pending_2", "running_1", "running_2", "new_lora"}
self.assertEqual(call_args, expected)
if __name__ == "__main__":
try: