fix: preserve priority for batched embedding requests (#32977)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
db8f3cdd11
commit
fe6a05a8e8
@@ -1161,6 +1161,7 @@ class EmbeddingReqInput:
|
||||
lora_id=self.lora_id[i] if self.lora_id is not None else None,
|
||||
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
|
||||
http_worker_ipc=self.http_worker_ipc,
|
||||
priority=self.priority,
|
||||
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
||||
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||
multi_item_delimiter_indices=(
|
||||
@@ -1188,6 +1189,7 @@ class EmbeddingReqInput:
|
||||
lora_id=self.lora_id[i] if self.lora_id is not None else None,
|
||||
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
|
||||
http_worker_ipc=self.http_worker_ipc,
|
||||
priority=self.priority,
|
||||
dimensions=self.dimensions,
|
||||
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
||||
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import copy
|
||||
import unittest
|
||||
|
||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
|
||||
from sglang.test.ci.ci_register import (
|
||||
register_amd_ci,
|
||||
register_cpu_ci,
|
||||
@@ -704,5 +704,25 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
|
||||
class TestEmbeddingReqInputGetItem(CustomTestCase):
|
||||
"""Test EmbeddingReqInput.__getitem__."""
|
||||
|
||||
def test_priority_is_preserved(self):
|
||||
"""Priority must survive the batch split, in both __getitem__ branches."""
|
||||
req = EmbeddingReqInput(text=["Hello", "World"], priority=7)
|
||||
req.normalize_batch_and_arguments()
|
||||
self.assertEqual([req[0].priority, req[1].priority], [7, 7])
|
||||
|
||||
cross_encoder_req = EmbeddingReqInput(
|
||||
text=[["query 1", "doc 1"], ["query 2", "doc 2"]],
|
||||
is_cross_encoder_request=True,
|
||||
priority=3,
|
||||
)
|
||||
cross_encoder_req.normalize_batch_and_arguments()
|
||||
self.assertEqual(
|
||||
[cross_encoder_req[0].priority, cross_encoder_req[1].priority], [3, 3]
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user