Feat/support rerank (#6058)

This commit is contained in:
woodx
2025-06-17 01:50:01 +08:00
committed by GitHub
parent 91a066ec6a
commit e30ef368ab
20 changed files with 684 additions and 30 deletions

View File

@@ -327,6 +327,20 @@ class Engine(EngineBase):
generator = self.tokenizer_manager.generate_request(obj, None)
return await generator.__anext__()
def rerank(
self,
prompt: Union[List[List[str]]],
) -> Dict:
"""
The arguments of this function is the same as `sglang/srt/managers/io_struct.py::EmbeddingReqInput`.
Please refer to `EmbeddingReqInput` for the documentation.
"""
obj = EmbeddingReqInput(text=prompt, is_cross_encoder_request=True)
loop = asyncio.get_event_loop()
generator = self.tokenizer_manager.generate_request(obj, None)
ret = loop.run_until_complete(generator.__anext__())
return ret
def shutdown(self):
"""Shutdown the engine"""
kill_process_tree(os.getpid(), include_parent=False)

View File

@@ -67,6 +67,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightFromDiskReqInput,
UpdateWeightsFromDistributedReqInput,
UpdateWeightsFromTensorReqInput,
V1RerankReqInput,
VertexGenerateReqInput,
)
from sglang.srt.managers.tokenizer_manager import TokenizerManager
@@ -79,6 +80,7 @@ from sglang.srt.openai_api.adapter import (
v1_delete_file,
v1_embeddings,
v1_files_create,
v1_rerank,
v1_retrieve_batch,
v1_retrieve_file,
v1_retrieve_file_content,
@@ -328,6 +330,15 @@ async def classify_request(obj: EmbeddingReqInput, request: Request):
return _create_error_response(e)
@app.api_route("/v1/rerank", methods=["POST", "PUT"])
async def v1_rerank_request(obj: V1RerankReqInput, raw_request: Request):
try:
ret = await v1_rerank(_global_state.tokenizer_manager, obj, raw_request)
return ret
except ValueError as e:
return _create_error_response(e)
@app.api_route("/flush_cache", methods=["GET", "POST"])
async def flush_cache():
"""Flush the radix cache."""