allow loras to be implicitly evicted and loaded based on max_loaded_loras (#11526)
This commit is contained in:
@@ -14,6 +14,7 @@
|
|||||||
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from collections import OrderedDict
|
||||||
from dataclasses import dataclass, field, fields
|
from dataclasses import dataclass, field, fields
|
||||||
from typing import Dict, List, Optional, Union
|
from typing import Dict, List, Optional, Union
|
||||||
from uuid import uuid4
|
from uuid import uuid4
|
||||||
@@ -71,8 +72,11 @@ class LoRARegistry:
|
|||||||
# Please note that the counter increment/decrement operations are not synchronized through this
|
# Please note that the counter increment/decrement operations are not synchronized through this
|
||||||
# lock, as they are designed to be non-blocking and can be performed concurrently.
|
# lock, as they are designed to be non-blocking and can be performed concurrently.
|
||||||
self._registry_lock = RWLock()
|
self._registry_lock = RWLock()
|
||||||
# A dictionary to hold LoRARef objects, mapping from LoRA name to LoRARef.
|
# An ordered dictionary to hold LoRARef objects, mapping from LoRA name to LoRARef.
|
||||||
self._registry: Dict[str, LoRARef] = {}
|
# The LoRARefs are stored in LRU order, such that LoRA adapters that have been
|
||||||
|
# most recently used are stored at the end. Note that lookups count for accesses.
|
||||||
|
# Ties are broken arbitrarily.
|
||||||
|
self._registry: OrderedDict[str, LoRARef] = OrderedDict()
|
||||||
# Counters for ongoing requests, mapping from LoRA ID to ConcurrentCounter.
|
# Counters for ongoing requests, mapping from LoRA ID to ConcurrentCounter.
|
||||||
self._counters: Dict[str, ConcurrentCounter] = {}
|
self._counters: Dict[str, ConcurrentCounter] = {}
|
||||||
|
|
||||||
@@ -124,29 +128,30 @@ class LoRARegistry:
|
|||||||
f"The following requested LoRA adapters are not loaded: {name}\n"
|
f"The following requested LoRA adapters are not loaded: {name}\n"
|
||||||
f"Loaded adapters: {self._registry.keys()}."
|
f"Loaded adapters: {self._registry.keys()}."
|
||||||
)
|
)
|
||||||
|
self._registry.move_to_end(name)
|
||||||
return lora_ref.lora_id
|
return lora_ref.lora_id
|
||||||
|
|
||||||
async with self._registry_lock.reader_lock:
|
if isinstance(lora_name, str):
|
||||||
if isinstance(lora_name, str):
|
async with self._registry_lock.writer_lock:
|
||||||
lora_id = _lookup(lora_name)
|
lora_id = _lookup(lora_name)
|
||||||
await self._counters[lora_id].increment(notify_all=False)
|
|
||||||
return lora_id
|
await self._counters[lora_id].increment(notify_all=False)
|
||||||
elif isinstance(lora_name, list):
|
return lora_id
|
||||||
|
elif isinstance(lora_name, list):
|
||||||
|
async with self._registry_lock.writer_lock:
|
||||||
lora_ids = [_lookup(name) for name in lora_name]
|
lora_ids = [_lookup(name) for name in lora_name]
|
||||||
|
|
||||||
# Increment the counters only after all IDs are looked up.
|
# Increment the counters only after all IDs are looked up.
|
||||||
await asyncio.gather(
|
await asyncio.gather(
|
||||||
*[
|
*[
|
||||||
self._counters[id].increment(notify_all=False)
|
self._counters[id].increment(notify_all=False)
|
||||||
for id in lora_ids
|
for id in lora_ids
|
||||||
if id is not None
|
if id is not None
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
return lora_ids
|
return lora_ids
|
||||||
else:
|
else:
|
||||||
raise TypeError(
|
raise TypeError("lora_name must be either a string or a list of strings.")
|
||||||
"lora_name must be either a string or a list of strings."
|
|
||||||
)
|
|
||||||
|
|
||||||
async def release(self, lora_id: Union[str, List[str]]):
|
async def release(self, lora_id: Union[str, List[str]]):
|
||||||
"""
|
"""
|
||||||
@@ -186,6 +191,37 @@ class LoRARegistry:
|
|||||||
await self._counters[lora_id].wait_for_zero()
|
await self._counters[lora_id].wait_for_zero()
|
||||||
del self._counters[lora_id]
|
del self._counters[lora_id]
|
||||||
|
|
||||||
|
async def get_unregistered_loras(self, lora_name: set[str]):
|
||||||
|
"""
|
||||||
|
Returns all LoRA adapters in lora_name that are not found in self._registry.
|
||||||
|
"""
|
||||||
|
async with self._registry_lock.writer_lock:
|
||||||
|
unregistered_loras = []
|
||||||
|
|
||||||
|
for name in lora_name:
|
||||||
|
if name in self._registry:
|
||||||
|
# This counts as a lookup, so we want to update the cache
|
||||||
|
self._registry.move_to_end(name)
|
||||||
|
else:
|
||||||
|
unregistered_loras.append(name)
|
||||||
|
|
||||||
|
return unregistered_loras
|
||||||
|
|
||||||
|
async def lru_lora_name(self, exclude_pinned=False):
|
||||||
|
"""
|
||||||
|
Returns the least recently used LoRA adapter.
|
||||||
|
If exclude_pinned is True, then return the LRU LoRA adapter that isn't pinned.
|
||||||
|
"""
|
||||||
|
async with self._registry_lock.reader_lock:
|
||||||
|
if not exclude_pinned:
|
||||||
|
return next(iter(self._registry), None)
|
||||||
|
|
||||||
|
for lora_name, lora_ref in self._registry.items():
|
||||||
|
if not lora_ref.pinned:
|
||||||
|
return lora_name
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
def _register_adapter(self, lora_ref: LoRARef):
|
def _register_adapter(self, lora_ref: LoRARef):
|
||||||
"""
|
"""
|
||||||
Internal helper method to register a LoRA adapter.
|
Internal helper method to register a LoRA adapter.
|
||||||
|
|||||||
@@ -495,6 +495,26 @@ class TokenizerCommunicatorMixin:
|
|||||||
|
|
||||||
return success, message
|
return success, message
|
||||||
|
|
||||||
|
async def _unload_lora_adapter_locked(
|
||||||
|
self: TokenizerManager,
|
||||||
|
obj: UnloadLoRAAdapterReqInput,
|
||||||
|
) -> UnloadLoRAAdapterReqOutput:
|
||||||
|
assert (
|
||||||
|
self.lora_update_lock.locked()
|
||||||
|
), "self.lora_update_lock must be locked in order for self._unload_lora_adapter_locked() to be called"
|
||||||
|
|
||||||
|
# Unregister the LoRA adapter from the registry to stop new requests for this adapter
|
||||||
|
# from being started.
|
||||||
|
lora_id = await self.lora_registry.unregister(obj.lora_name)
|
||||||
|
obj.lora_id = lora_id
|
||||||
|
|
||||||
|
# Initiate the actual unloading operation at the backend processes only after all
|
||||||
|
# ongoing requests using this LoRA adapter are finished.
|
||||||
|
await self.lora_registry.wait_for_unload(lora_id)
|
||||||
|
result = (await self.update_lora_adapter_communicator(obj))[0]
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
async def load_lora_adapter(
|
async def load_lora_adapter(
|
||||||
self: TokenizerManager,
|
self: TokenizerManager,
|
||||||
obj: LoadLoRAAdapterReqInput,
|
obj: LoadLoRAAdapterReqInput,
|
||||||
@@ -520,17 +540,6 @@ class TokenizerCommunicatorMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
async with self.lora_update_lock:
|
async with self.lora_update_lock:
|
||||||
if (
|
|
||||||
self.server_args.max_loaded_loras is not None
|
|
||||||
and self.lora_registry.num_registered_loras
|
|
||||||
>= self.server_args.max_loaded_loras
|
|
||||||
):
|
|
||||||
raise ValueError(
|
|
||||||
f"Cannot load LoRA adapter {obj.lora_name} at path {obj.lora_path}. "
|
|
||||||
f"Maximum number of loaded LoRA adapters is {self.server_args.max_loaded_loras}. "
|
|
||||||
"Please unload some LoRA adapters before loading new ones."
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate new uniquely identifiable LoRARef object.
|
# Generate new uniquely identifiable LoRARef object.
|
||||||
new_adapter = LoRARef(
|
new_adapter = LoRARef(
|
||||||
lora_name=obj.lora_name,
|
lora_name=obj.lora_name,
|
||||||
@@ -545,6 +554,37 @@ class TokenizerCommunicatorMixin:
|
|||||||
# Register the LoRA adapter only after loading is successful.
|
# Register the LoRA adapter only after loading is successful.
|
||||||
if result.success:
|
if result.success:
|
||||||
await self.lora_registry.register(new_adapter)
|
await self.lora_registry.register(new_adapter)
|
||||||
|
self.lora_ref_cache[obj.lora_name] = new_adapter
|
||||||
|
|
||||||
|
if self.server_args.max_loaded_loras is not None:
|
||||||
|
while (
|
||||||
|
self.lora_registry.num_registered_loras
|
||||||
|
> self.server_args.max_loaded_loras
|
||||||
|
):
|
||||||
|
lru_lora_name = await self.lora_registry.lru_lora_name(
|
||||||
|
exclude_pinned=True
|
||||||
|
)
|
||||||
|
if lru_lora_name is None:
|
||||||
|
raise ValueError(
|
||||||
|
"Didn't find any LoRA adapters when trying to evict LRU LoRA adapter. "
|
||||||
|
f"LoRA registry is: {self.lora_registry._registry}"
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
|
||||||
|
f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
|
||||||
|
f"max allowed: {self.server_args.max_loaded_loras})"
|
||||||
|
)
|
||||||
|
|
||||||
|
unload_result = await self._unload_lora_adapter_locked(
|
||||||
|
UnloadLoRAAdapterReqInput(lora_name=lru_lora_name)
|
||||||
|
)
|
||||||
|
if not unload_result.success:
|
||||||
|
raise ValueError(
|
||||||
|
f"Error while unloading LRU LoRA adapter '{lru_lora_name}': "
|
||||||
|
f"{unload_result.error_message}"
|
||||||
|
)
|
||||||
|
del result.loaded_adapters[lru_lora_name]
|
||||||
|
|
||||||
return result
|
return result
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
@@ -581,17 +621,7 @@ class TokenizerCommunicatorMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
async with self.lora_update_lock:
|
async with self.lora_update_lock:
|
||||||
# Unregister the LoRA adapter from the registry to stop new requests for this adapter
|
return await self._unload_lora_adapter_locked(obj)
|
||||||
# from being started.
|
|
||||||
lora_id = await self.lora_registry.unregister(obj.lora_name)
|
|
||||||
obj.lora_id = lora_id
|
|
||||||
|
|
||||||
# Initiate the actual unloading operation at the backend processes only after all
|
|
||||||
# ongoing requests using this LoRA adapter are finished.
|
|
||||||
await self.lora_registry.wait_for_unload(lora_id)
|
|
||||||
result = (await self.update_lora_adapter_communicator(obj))[0]
|
|
||||||
|
|
||||||
return result
|
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
return UnloadLoRAAdapterReqOutput(success=False, error_message=str(e))
|
return UnloadLoRAAdapterReqOutput(success=False, error_message=str(e))
|
||||||
|
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ from fastapi import BackgroundTasks
|
|||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.lora.lora_registry import LoRARegistry
|
from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry
|
||||||
from sglang.srt.managers.async_dynamic_batch_tokenizer import AsyncDynamicbatchTokenizer
|
from sglang.srt.managers.async_dynamic_batch_tokenizer import AsyncDynamicbatchTokenizer
|
||||||
from sglang.srt.managers.async_mm_data_processor import AsyncMMDataProcessor
|
from sglang.srt.managers.async_mm_data_processor import AsyncMMDataProcessor
|
||||||
from sglang.srt.managers.disagg_service import start_disagg_service
|
from sglang.srt.managers.disagg_service import start_disagg_service
|
||||||
@@ -60,6 +60,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
GenerateReqInput,
|
GenerateReqInput,
|
||||||
GetLoadReqInput,
|
GetLoadReqInput,
|
||||||
HealthCheckOutput,
|
HealthCheckOutput,
|
||||||
|
LoadLoRAAdapterReqInput,
|
||||||
OpenSessionReqOutput,
|
OpenSessionReqOutput,
|
||||||
SessionParams,
|
SessionParams,
|
||||||
TokenizedEmbeddingReqInput,
|
TokenizedEmbeddingReqInput,
|
||||||
@@ -357,6 +358,13 @@ class TokenizerManager(TokenizerCommunicatorMixin):
|
|||||||
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
|
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
|
||||||
# LoRA updates and inference to overlap.
|
# LoRA updates and inference to overlap.
|
||||||
self.lora_update_lock = asyncio.Lock()
|
self.lora_update_lock = asyncio.Lock()
|
||||||
|
# A cache for mapping the lora_name for LoRA adapters that have been loaded at any
|
||||||
|
# point to their latest LoRARef objects, so that they can be
|
||||||
|
# dynamically loaded if needed for inference
|
||||||
|
self.lora_ref_cache: Dict[str, LoRARef] = {}
|
||||||
|
if self.server_args.lora_paths is not None:
|
||||||
|
for lora_ref in self.server_args.lora_paths:
|
||||||
|
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
|
||||||
|
|
||||||
# Disaggregation
|
# Disaggregation
|
||||||
self.disaggregation_mode = DisaggregationMode(
|
self.disaggregation_mode = DisaggregationMode(
|
||||||
@@ -448,6 +456,51 @@ class TokenizerManager(TokenizerCommunicatorMixin):
|
|||||||
|
|
||||||
async with self.model_update_lock.reader_lock:
|
async with self.model_update_lock.reader_lock:
|
||||||
if self.server_args.enable_lora and obj.lora_path:
|
if self.server_args.enable_lora and obj.lora_path:
|
||||||
|
if isinstance(obj.lora_path, str):
|
||||||
|
unique_lora_paths = set([obj.lora_path])
|
||||||
|
else:
|
||||||
|
unique_lora_paths = set(obj.lora_path)
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.server_args.max_loaded_loras is not None
|
||||||
|
and len(unique_lora_paths) > self.server_args.max_loaded_loras
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"Received request with {len(unique_lora_paths)} unique loras requested "
|
||||||
|
f"but max loaded loras is {self.server_args.max_loaded_loras}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Reload all existing LoRA adapters that have been dynamically unloaded
|
||||||
|
unregistered_loras = await self.lora_registry.get_unregistered_loras(
|
||||||
|
unique_lora_paths
|
||||||
|
)
|
||||||
|
for lora_path in unregistered_loras:
|
||||||
|
if lora_path is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if lora_path not in self.lora_ref_cache:
|
||||||
|
raise ValueError(
|
||||||
|
f"Got LoRA adapter that has never been loaded: {lora_path}\n"
|
||||||
|
f"All loaded adapters: {self.lora_ref_cache.keys()}."
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(f"Reloading evicted adapter: {lora_path}")
|
||||||
|
new_lora_ref = self.lora_ref_cache[lora_path]
|
||||||
|
load_result = await self.load_lora_adapter(
|
||||||
|
LoadLoRAAdapterReqInput(
|
||||||
|
lora_name=new_lora_ref.lora_name,
|
||||||
|
lora_path=new_lora_ref.lora_path,
|
||||||
|
pinned=new_lora_ref.pinned,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
not load_result.success
|
||||||
|
and "already loaded" not in load_result.error_message
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"Failed to implicitly load LoRA adapter {lora_path}: {load_result.error_message}"
|
||||||
|
)
|
||||||
|
|
||||||
# Look up the LoRA ID from the registry and start tracking ongoing LoRA requests.
|
# Look up the LoRA ID from the registry and start tracking ongoing LoRA requests.
|
||||||
obj.lora_id = await self.lora_registry.acquire(obj.lora_path)
|
obj.lora_id = await self.lora_registry.acquire(obj.lora_path)
|
||||||
|
|
||||||
|
|||||||
@@ -58,6 +58,9 @@ class Operation:
|
|||||||
data: Optional[Any]
|
data: Optional[Any]
|
||||||
# If the operation is expected to fail, this is the error message to expect
|
# If the operation is expected to fail, this is the error message to expect
|
||||||
expected_error: Optional[str] = None
|
expected_error: Optional[str] = None
|
||||||
|
# Because the logic for implicitly evicting LoRA adapters can be complicated, we explicitly
|
||||||
|
# pass in LoRA adapters that should be implicitly evicted here
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -142,20 +145,20 @@ BASIC_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
expected_error="already loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="pbevan11/llama-3.1-8b-ocr-correction",
|
data="pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
expected_error="already loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -174,7 +177,6 @@ BASIC_TESTS = [
|
|||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -198,12 +200,10 @@ BASIC_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -269,7 +269,6 @@ BASIC_TESTS = [
|
|||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -295,12 +294,10 @@ BASIC_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
||||||
expected_error="not loaded",
|
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -309,14 +306,17 @@ BASIC_TESTS = [
|
|||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
data="philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
expected_error="already loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
expected_error="already loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="pbevan11/llama-3.1-8b-ocr-correction",
|
data="pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
expected_error="already loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -364,7 +364,7 @@ TARGET_MODULE_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -404,7 +404,7 @@ TARGET_MODULE_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"algoprog/fact-generation-llama-3.1-8b-instruct-lora"
|
"algoprog/fact-generation-llama-3.1-8b-instruct-lora"
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -444,7 +444,7 @@ TARGET_MODULE_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -485,12 +485,12 @@ MAX_LORA_RANK_TESTS = [
|
|||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
data=create_batch_data("pbevan11/llama-3.1-8b-ocr-correction"),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -516,7 +516,7 @@ MAX_LORA_RANK_TESTS = [
|
|||||||
data=create_batch_data(
|
data=create_batch_data(
|
||||||
"philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
),
|
),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
@@ -553,7 +553,7 @@ MAX_LORA_RANK_TESTS = [
|
|||||||
Operation(
|
Operation(
|
||||||
type=OperationType.FORWARD,
|
type=OperationType.FORWARD,
|
||||||
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
data=create_batch_data("philschmid/code-llama-3-1-8b-text-to-sql-lora"),
|
||||||
expected_error="not loaded",
|
expected_error="never been loaded",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -580,7 +580,7 @@ MAX_LORA_RANK_TESTS = [
|
|||||||
]
|
]
|
||||||
MAX_LOADED_LORAS_TESTS = [
|
MAX_LOADED_LORAS_TESTS = [
|
||||||
TestCase(
|
TestCase(
|
||||||
description="Test max_loaded_loras limit",
|
description="Test max_loaded_loras limit as well as implicit eviction and reloading",
|
||||||
base="meta-llama/Llama-3.1-8B-Instruct",
|
base="meta-llama/Llama-3.1-8B-Instruct",
|
||||||
max_loras_per_batch=2,
|
max_loras_per_batch=2,
|
||||||
max_loaded_loras=2,
|
max_loaded_loras=2,
|
||||||
@@ -598,15 +598,143 @@ MAX_LOADED_LORAS_TESTS = [
|
|||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="pbevan11/llama-3.1-8b-ocr-correction",
|
data="pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
expected_error="Maximum number of loaded LoRA adapters",
|
expected_implicit_evictions={
|
||||||
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora"
|
||||||
|
},
|
||||||
|
),
|
||||||
|
# Implicitly load "philschmid/code-llama-3-1-8b-text-to-sql-lora"
|
||||||
|
Operation(
|
||||||
|
type=OperationType.FORWARD,
|
||||||
|
data=create_batch_data(
|
||||||
|
[
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
]
|
||||||
|
),
|
||||||
|
expected_implicit_evictions={"pbevan11/llama-3.1-8b-ocr-correction"},
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.UNLOAD,
|
||||||
|
data="philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.FORWARD,
|
||||||
|
data=create_batch_data(
|
||||||
|
[
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
]
|
||||||
|
),
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.LOAD,
|
||||||
|
data="philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
),
|
||||||
|
# Implicitly load "pbevan11/llama-3.1-8b-ocr-correction" and make sure that "Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
|
# isn't implicitly unloaded even though it is LRU because it is needed for this forward pass
|
||||||
|
Operation(
|
||||||
|
type=OperationType.FORWARD,
|
||||||
|
data=create_batch_data(
|
||||||
|
[
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
"pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
]
|
||||||
|
),
|
||||||
|
expected_implicit_evictions={
|
||||||
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora"
|
||||||
|
},
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.UNLOAD,
|
type=OperationType.UNLOAD,
|
||||||
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
),
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.LOAD,
|
||||||
|
data="algoprog/fact-generation-llama-3.1-8b-instruct-lora",
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.FORWARD,
|
||||||
|
data=create_batch_data(
|
||||||
|
[
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
]
|
||||||
|
),
|
||||||
|
expected_implicit_evictions={
|
||||||
|
"pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
"algoprog/fact-generation-llama-3.1-8b-instruct-lora",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
TestCase(
|
||||||
|
description="Test implicit eviction and reloading with pinned LoRA adapters",
|
||||||
|
base="meta-llama/Llama-3.1-8B-Instruct",
|
||||||
|
max_loras_per_batch=2,
|
||||||
|
max_loaded_loras=2,
|
||||||
|
all_adapters=[
|
||||||
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
"pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
],
|
||||||
|
initial_adapters=[
|
||||||
|
{
|
||||||
|
"lora_name": "philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
"lora_path": "philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
"pinned": True,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
op_sequence=[
|
||||||
|
Operation(
|
||||||
|
type=OperationType.LOAD,
|
||||||
|
data="Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
data="pbevan11/llama-3.1-8b-ocr-correction",
|
data="pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
expected_implicit_evictions={
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
|
},
|
||||||
|
),
|
||||||
|
# Implicitly load "Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16"
|
||||||
|
Operation(
|
||||||
|
type=OperationType.FORWARD,
|
||||||
|
data=create_batch_data(
|
||||||
|
[
|
||||||
|
"philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
]
|
||||||
|
),
|
||||||
|
expected_implicit_evictions={"pbevan11/llama-3.1-8b-ocr-correction"},
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.LOAD,
|
||||||
|
data={
|
||||||
|
"lora_name": "pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
"lora_path": "pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
"pinned": True,
|
||||||
|
},
|
||||||
|
expected_error="unpin some adapters",
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.UNLOAD,
|
||||||
|
data="philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.LOAD,
|
||||||
|
data={
|
||||||
|
"lora_name": "pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
"lora_path": "pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
"pinned": True,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
Operation(
|
||||||
|
type=OperationType.FORWARD,
|
||||||
|
data=create_batch_data(
|
||||||
|
[
|
||||||
|
"Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
|
"pbevan11/llama-3.1-8b-ocr-correction",
|
||||||
|
]
|
||||||
|
),
|
||||||
),
|
),
|
||||||
],
|
],
|
||||||
),
|
),
|
||||||
@@ -640,7 +768,7 @@ EVICTION_TESTS = [
|
|||||||
"lora_path": "Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
"lora_path": "Nutanix/Meta-Llama-3.1-8B-Instruct_lora_4_alpha_16",
|
||||||
"pinned": True,
|
"pinned": True,
|
||||||
},
|
},
|
||||||
expected_error="starvation",
|
expected_error="unpin some adapters",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -681,7 +809,7 @@ EVICTION_TESTS = [
|
|||||||
"lora_path": "philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
"lora_path": "philschmid/code-llama-3-1-8b-text-to-sql-lora",
|
||||||
"pinned": True,
|
"pinned": True,
|
||||||
},
|
},
|
||||||
expected_error="starvation",
|
expected_error="unpin some adapters",
|
||||||
),
|
),
|
||||||
Operation(
|
Operation(
|
||||||
type=OperationType.LOAD,
|
type=OperationType.LOAD,
|
||||||
@@ -805,6 +933,7 @@ class LoRAUpdateTestSessionBase:
|
|||||||
lora_name: str,
|
lora_name: str,
|
||||||
lora_path: Optional[str] = None,
|
lora_path: Optional[str] = None,
|
||||||
expected_error: Optional[str] = None,
|
expected_error: Optional[str] = None,
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Load a LoRA adapter by name and path.
|
Load a LoRA adapter by name and path.
|
||||||
@@ -823,6 +952,7 @@ class LoRAUpdateTestSessionBase:
|
|||||||
lora_paths: List[str],
|
lora_paths: List[str],
|
||||||
max_new_tokens: int = 32,
|
max_new_tokens: int = 32,
|
||||||
expected_error: Optional[str] = None,
|
expected_error: Optional[str] = None,
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Perform a batch forward pass with the current set of loaded LoRA adapters.
|
Perform a batch forward pass with the current set of loaded LoRA adapters.
|
||||||
@@ -869,6 +999,7 @@ class LoRAUpdateEngineTestSession(LoRAUpdateTestSessionBase):
|
|||||||
lora_path: Optional[str] = None,
|
lora_path: Optional[str] = None,
|
||||||
expected_error: Optional[str] = None,
|
expected_error: Optional[str] = None,
|
||||||
pinned: bool = False,
|
pinned: bool = False,
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Load a LoRA adapter by name and path.
|
Load a LoRA adapter by name and path.
|
||||||
@@ -893,6 +1024,9 @@ class LoRAUpdateEngineTestSession(LoRAUpdateTestSessionBase):
|
|||||||
print(f"Received error as expected: {response.error_message}")
|
print(f"Received error as expected: {response.error_message}")
|
||||||
else:
|
else:
|
||||||
self.expected_adapters.add(lora_name)
|
self.expected_adapters.add(lora_name)
|
||||||
|
if expected_implicit_evictions is not None:
|
||||||
|
self.expected_adapters -= expected_implicit_evictions
|
||||||
|
|
||||||
self.testcase.assertTrue(
|
self.testcase.assertTrue(
|
||||||
response.success,
|
response.success,
|
||||||
f"Failed to load LoRA adapter {lora_name}: {response.error_message}",
|
f"Failed to load LoRA adapter {lora_name}: {response.error_message}",
|
||||||
@@ -933,6 +1067,7 @@ class LoRAUpdateEngineTestSession(LoRAUpdateTestSessionBase):
|
|||||||
lora_paths: List[str],
|
lora_paths: List[str],
|
||||||
max_new_tokens: int = 32,
|
max_new_tokens: int = 32,
|
||||||
expected_error: Optional[str] = None,
|
expected_error: Optional[str] = None,
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Perform a batch forward pass with the current set of loaded LoRA adapters.
|
Perform a batch forward pass with the current set of loaded LoRA adapters.
|
||||||
@@ -964,6 +1099,13 @@ class LoRAUpdateEngineTestSession(LoRAUpdateTestSessionBase):
|
|||||||
output = response.output_strs
|
output = response.output_strs
|
||||||
print(f"output_strs: {output}")
|
print(f"output_strs: {output}")
|
||||||
|
|
||||||
|
self.expected_adapters.update(
|
||||||
|
[lora_path for lora_path in lora_paths if lora_path is not None]
|
||||||
|
)
|
||||||
|
|
||||||
|
if expected_implicit_evictions is not None:
|
||||||
|
self.expected_adapters -= expected_implicit_evictions
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
@@ -1026,6 +1168,7 @@ class LoRAUpdateServerTestSession(LoRAUpdateTestSessionBase):
|
|||||||
lora_path: Optional[str] = None,
|
lora_path: Optional[str] = None,
|
||||||
expected_error: Optional[str] = None,
|
expected_error: Optional[str] = None,
|
||||||
pinned: bool = False,
|
pinned: bool = False,
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Load a LoRA adapter by name and path.
|
Load a LoRA adapter by name and path.
|
||||||
@@ -1051,6 +1194,9 @@ class LoRAUpdateServerTestSession(LoRAUpdateTestSessionBase):
|
|||||||
print(f"Received error as expected: {response.text}")
|
print(f"Received error as expected: {response.text}")
|
||||||
else:
|
else:
|
||||||
self.expected_adapters.add(lora_name)
|
self.expected_adapters.add(lora_name)
|
||||||
|
if expected_implicit_evictions is not None:
|
||||||
|
self.expected_adapters -= expected_implicit_evictions
|
||||||
|
|
||||||
self.testcase.assertTrue(
|
self.testcase.assertTrue(
|
||||||
response.ok, f"Failed to load LoRA adapter {lora_name}: {response.text}"
|
response.ok, f"Failed to load LoRA adapter {lora_name}: {response.text}"
|
||||||
)
|
)
|
||||||
@@ -1072,6 +1218,7 @@ class LoRAUpdateServerTestSession(LoRAUpdateTestSessionBase):
|
|||||||
DEFAULT_URL_FOR_TEST + "/unload_lora_adapter",
|
DEFAULT_URL_FOR_TEST + "/unload_lora_adapter",
|
||||||
json={"lora_name": lora_name},
|
json={"lora_name": lora_name},
|
||||||
)
|
)
|
||||||
|
|
||||||
self.testcase.assertTrue(
|
self.testcase.assertTrue(
|
||||||
response.ok, f"Failed to unload LoRA adapter {lora_name}: {response.text}"
|
response.ok, f"Failed to unload LoRA adapter {lora_name}: {response.text}"
|
||||||
)
|
)
|
||||||
@@ -1090,6 +1237,7 @@ class LoRAUpdateServerTestSession(LoRAUpdateTestSessionBase):
|
|||||||
lora_paths: List[str],
|
lora_paths: List[str],
|
||||||
max_new_tokens: int = 32,
|
max_new_tokens: int = 32,
|
||||||
expected_error: Optional[str] = None,
|
expected_error: Optional[str] = None,
|
||||||
|
expected_implicit_evictions: Optional[set[str]] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Perform a batch forward pass with the current set of loaded LoRA adapters.
|
Perform a batch forward pass with the current set of loaded LoRA adapters.
|
||||||
@@ -1131,6 +1279,14 @@ class LoRAUpdateServerTestSession(LoRAUpdateTestSessionBase):
|
|||||||
f"Expected {len(prompts)} outputs, but got {len(output)}",
|
f"Expected {len(prompts)} outputs, but got {len(output)}",
|
||||||
)
|
)
|
||||||
print(f"output_strs: {output}")
|
print(f"output_strs: {output}")
|
||||||
|
|
||||||
|
self.expected_adapters.update(
|
||||||
|
[lora_path for lora_path in lora_paths if lora_path is not None]
|
||||||
|
)
|
||||||
|
|
||||||
|
if expected_implicit_evictions is not None:
|
||||||
|
self.expected_adapters -= expected_implicit_evictions
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
@@ -1192,6 +1348,7 @@ class TestLoRADynamicUpdate(CustomTestCase):
|
|||||||
op_type = op.type
|
op_type = op.type
|
||||||
data = op.data
|
data = op.data
|
||||||
expected_error = op.expected_error
|
expected_error = op.expected_error
|
||||||
|
expected_implicit_evictions = op.expected_implicit_evictions
|
||||||
print("-" * 100)
|
print("-" * 100)
|
||||||
print(
|
print(
|
||||||
f"Running operation: {op_type} --- data: {data} --- mode: {mode} ---"
|
f"Running operation: {op_type} --- data: {data} --- mode: {mode} ---"
|
||||||
@@ -1208,6 +1365,7 @@ class TestLoRADynamicUpdate(CustomTestCase):
|
|||||||
|
|
||||||
result = session.load_lora_adapter(
|
result = session.load_lora_adapter(
|
||||||
expected_error=expected_error,
|
expected_error=expected_error,
|
||||||
|
expected_implicit_evictions=expected_implicit_evictions,
|
||||||
**adapter_info,
|
**adapter_info,
|
||||||
)
|
)
|
||||||
elif op_type == OperationType.UNLOAD:
|
elif op_type == OperationType.UNLOAD:
|
||||||
@@ -1221,6 +1379,7 @@ class TestLoRADynamicUpdate(CustomTestCase):
|
|||||||
lora_paths=list(adapters),
|
lora_paths=list(adapters),
|
||||||
max_new_tokens=max_new_tokens,
|
max_new_tokens=max_new_tokens,
|
||||||
expected_error=expected_error,
|
expected_error=expected_error,
|
||||||
|
expected_implicit_evictions=expected_implicit_evictions,
|
||||||
)
|
)
|
||||||
if not expected_error:
|
if not expected_error:
|
||||||
forward_outputs.append(result)
|
forward_outputs.append(result)
|
||||||
|
|||||||
Reference in New Issue
Block a user