add support to enable lora with embedding models (#17780)
Co-authored-by: Vedant Jhaveri <vjhaveri@linkedin.com>
This commit is contained in:
@@ -379,6 +379,7 @@ class Engine(EngineBase):
|
||||
audio_data: Optional[MultimodalDataInputFormat] = None,
|
||||
video_data: Optional[MultimodalDataInputFormat] = None,
|
||||
dimensions: Optional[int] = None,
|
||||
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None,
|
||||
external_trace_header: Optional[Dict] = None,
|
||||
rid: Optional[Union[List[str], str]] = None,
|
||||
) -> Dict:
|
||||
@@ -392,6 +393,7 @@ class Engine(EngineBase):
|
||||
audio_data=audio_data,
|
||||
video_data=video_data,
|
||||
dimensions=dimensions,
|
||||
lora_path=lora_path,
|
||||
external_trace_header=external_trace_header,
|
||||
rid=rid,
|
||||
)
|
||||
@@ -406,6 +408,7 @@ class Engine(EngineBase):
|
||||
audio_data: Optional[MultimodalDataInputFormat] = None,
|
||||
video_data: Optional[MultimodalDataInputFormat] = None,
|
||||
dimensions: Optional[int] = None,
|
||||
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None,
|
||||
external_trace_header: Optional[Dict] = None,
|
||||
rid: Optional[Union[List[str], str]] = None,
|
||||
) -> Dict:
|
||||
@@ -421,6 +424,7 @@ class Engine(EngineBase):
|
||||
audio_data=audio_data,
|
||||
video_data=video_data,
|
||||
dimensions=dimensions,
|
||||
lora_path=lora_path,
|
||||
external_trace_header=external_trace_header,
|
||||
rid=rid,
|
||||
)
|
||||
|
||||
@@ -897,6 +897,8 @@ class EmbeddingRequest(BaseModel):
|
||||
rid: Optional[Union[List[str], str]] = None
|
||||
# Priority for the request
|
||||
priority: Optional[int] = None
|
||||
# LoRA adapter path(s)
|
||||
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
|
||||
|
||||
|
||||
class EmbeddingObject(BaseModel):
|
||||
|
||||
@@ -126,12 +126,24 @@ class OpenAIServingEmbedding(OpenAIServingBase):
|
||||
# Other types (should not happen but handle gracefully)
|
||||
prompt_kwargs = {"input_ids": prompt}
|
||||
|
||||
# Resolve LoRA adapter from model parameter or explicit lora_path
|
||||
lora_path = self._resolve_lora_path(request.model, request.lora_path)
|
||||
if lora_path:
|
||||
first_adapter = (
|
||||
lora_path
|
||||
if isinstance(lora_path, str)
|
||||
else next((a for a in lora_path if a), None)
|
||||
)
|
||||
if first_adapter:
|
||||
self._validate_lora_enabled(first_adapter)
|
||||
|
||||
adapted_request = EmbeddingReqInput(
|
||||
**prompt_kwargs,
|
||||
rid=request.rid,
|
||||
priority=request.priority,
|
||||
routing_key=self.extract_routing_key(raw_request),
|
||||
dimensions=request.dimensions,
|
||||
lora_path=lora_path,
|
||||
)
|
||||
|
||||
return adapted_request, request
|
||||
|
||||
@@ -824,6 +824,11 @@ class EmbeddingReqInput(BaseReq, APIServingTimingMixin):
|
||||
# The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings.
|
||||
dimensions: Optional[int] = None
|
||||
|
||||
# The path to the LoRA adaptors
|
||||
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
|
||||
# The uid of LoRA adaptors, should be initialized by tokenizer manager
|
||||
lora_id: Optional[Union[List[Optional[str]], Optional[str]]] = None
|
||||
|
||||
def normalize_batch_and_arguments(self):
|
||||
# at least one of text, input_ids, or image should be provided
|
||||
if self.text is None and self.input_ids is None and self.image_data is None:
|
||||
@@ -875,6 +880,21 @@ class EmbeddingReqInput(BaseReq, APIServingTimingMixin):
|
||||
for i in range(self.batch_size):
|
||||
self.sampling_params[i]["max_new_tokens"] = 0
|
||||
|
||||
self._normalize_lora_paths(self.batch_size)
|
||||
|
||||
def _normalize_lora_paths(self, num):
|
||||
"""Normalize LoRA paths for batch processing."""
|
||||
if self.lora_path is not None:
|
||||
if isinstance(self.lora_path, str):
|
||||
self.lora_path = [self.lora_path] * num
|
||||
elif isinstance(self.lora_path, list):
|
||||
if len(self.lora_path) != num:
|
||||
raise ValueError(
|
||||
f"lora_path list length ({len(self.lora_path)}) must match batch size ({num})"
|
||||
)
|
||||
else:
|
||||
raise ValueError("lora_path should be a list or a string.")
|
||||
|
||||
def contains_mm_input(self) -> bool:
|
||||
return (
|
||||
has_valid_data(self.image_data)
|
||||
@@ -888,6 +908,8 @@ class EmbeddingReqInput(BaseReq, APIServingTimingMixin):
|
||||
text=[self.text[i]] if self.text is not None else None,
|
||||
sampling_params=self.sampling_params[i],
|
||||
rid=self.rid[i],
|
||||
lora_path=self.lora_path[i] if self.lora_path is not None else None,
|
||||
lora_id=self.lora_id[i] if self.lora_id is not None else None,
|
||||
is_cross_encoder_request=True,
|
||||
http_worker_ipc=self.http_worker_ipc,
|
||||
)
|
||||
@@ -900,6 +922,8 @@ class EmbeddingReqInput(BaseReq, APIServingTimingMixin):
|
||||
video_data=self.video_data[i] if self.video_data is not None else None,
|
||||
sampling_params=self.sampling_params[i],
|
||||
rid=self.rid[i],
|
||||
lora_path=self.lora_path[i] if self.lora_path is not None else None,
|
||||
lora_id=self.lora_id[i] if self.lora_id is not None else None,
|
||||
external_trace_header=self.external_trace_header,
|
||||
dimensions=self.dimensions,
|
||||
http_worker_ipc=self.http_worker_ipc,
|
||||
@@ -928,6 +952,8 @@ class TokenizedEmbeddingReqInput(BaseReq):
|
||||
priority: Optional[int] = None
|
||||
# The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings.
|
||||
dimensions: Optional[int] = None
|
||||
# LoRA related
|
||||
lora_id: Optional[str] = None # None means just use the base model
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1762,6 +1762,7 @@ class Scheduler(
|
||||
token_type_ids=recv_req.token_type_ids,
|
||||
priority=recv_req.priority,
|
||||
dimensions=recv_req.dimensions,
|
||||
lora_id=recv_req.lora_id,
|
||||
http_worker_ipc=recv_req.http_worker_ipc,
|
||||
)
|
||||
req.tokenizer = self.tokenizer
|
||||
|
||||
@@ -958,6 +958,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
rid=obj.rid,
|
||||
priority=obj.priority,
|
||||
dimensions=obj.dimensions,
|
||||
lora_id=obj.lora_id,
|
||||
http_worker_ipc=obj.http_worker_ipc,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user