fix: preserve priority for batched embedding requests (#32977)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Nikhil Kulkarni
2026-08-06 19:26:18 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent db8f3cdd11
commit fe6a05a8e8
2 changed files with 23 additions and 1 deletions
+2
View File
@@ -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()