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,
|
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),
|
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
|
||||||
http_worker_ipc=self.http_worker_ipc,
|
http_worker_ipc=self.http_worker_ipc,
|
||||||
|
priority=self.priority,
|
||||||
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
||||||
return_prompt_token_ids=self.return_prompt_token_ids,
|
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||||
multi_item_delimiter_indices=(
|
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,
|
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),
|
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
|
||||||
http_worker_ipc=self.http_worker_ipc,
|
http_worker_ipc=self.http_worker_ipc,
|
||||||
|
priority=self.priority,
|
||||||
dimensions=self.dimensions,
|
dimensions=self.dimensions,
|
||||||
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
||||||
return_prompt_token_ids=self.return_prompt_token_ids,
|
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import copy
|
import copy
|
||||||
import unittest
|
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 (
|
from sglang.test.ci.ci_register import (
|
||||||
register_amd_ci,
|
register_amd_ci,
|
||||||
register_cpu_ci,
|
register_cpu_ci,
|
||||||
@@ -704,5 +704,25 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
|||||||
req.normalize_batch_and_arguments()
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user