4b86b2a660
- server-dev start/stop/deploy 및 Gitea push 자동 배포 - local-dev 로컬 개발 환경 Co-authored-by: Cursor <cursoragent@cursor.com>
80 lines
3.0 KiB
Python
80 lines
3.0 KiB
Python
import math
|
|
import os
|
|
import torch
|
|
from fastapi import FastAPI, HTTPException
|
|
from pydantic import BaseModel
|
|
from typing import List
|
|
from vllm import LLM, SamplingParams
|
|
from vllm.inputs.data import TokensPrompt
|
|
from transformers import AutoTokenizer
|
|
|
|
# 오프라인 환경 변수 강제 설정
|
|
os.environ["HF_HUB_OFFLINE"] = "1"
|
|
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
|
|
|
app = FastAPI()
|
|
|
|
# 1. 모델 경로 (컨테이너 내부 경로 기준)
|
|
MODEL_PATH = "/data/qwen3-reranker-4b"
|
|
|
|
# 토크나이저 및 모델 로드
|
|
# vllm-openai 이미지에는 이미 transformers, vllm, fastapi가 들어있습니다.
|
|
tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH, local_files_only=True)
|
|
model = LLM(
|
|
model=MODEL_PATH,
|
|
gpu_memory_utilization=0.10, # H100 80GB 중 24GB 점유 (남은 공간은 80B 모델용)
|
|
max_model_len=4096,
|
|
max_num_seqs=20,
|
|
trust_remote_code=True,
|
|
dtype="float16",
|
|
enforce_eager=True # 오프라인 환경에서 불필요한 커널 컴파일 방지
|
|
)
|
|
|
|
# 토큰 설정
|
|
true_token = tokenizer("yes", add_special_tokens=False).input_ids[0]
|
|
false_token = tokenizer("no", add_special_tokens=False).input_ids[0]
|
|
suffix = "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"
|
|
suffix_tokens = tokenizer.encode(suffix, add_special_tokens=False)
|
|
|
|
class RerankRequest(BaseModel):
|
|
query: str
|
|
documents: List[str]
|
|
|
|
@app.post("/rerank")
|
|
async def rerank(request: RerankRequest):
|
|
task = 'Given a web search query, retrieve relevant passages that answer the query'
|
|
prompts = []
|
|
|
|
for doc in request.documents:
|
|
messages = [
|
|
{"role": "system", "content": "Judge whether the Document meets the requirements based on the Query and the Instruct provided. Note that the answer can only be \"yes\" or \"no\"."},
|
|
{"role": "user", "content": f"<Instruct>: {task}\n\n<Query>: {request.query}\n\n<Document>: {doc}"}
|
|
]
|
|
token_ids = tokenizer.apply_chat_template(messages, tokenize=True, add_generation_prompt=False)
|
|
# 길이 제한 및 suffix 추가
|
|
token_ids = token_ids[:8192 - len(suffix_tokens)] + suffix_tokens
|
|
prompts.append(TokensPrompt(prompt_token_ids=token_ids))
|
|
|
|
sampling_params = SamplingParams(
|
|
temperature=0, max_tokens=1, logprobs=20,
|
|
allowed_token_ids=[true_token, false_token]
|
|
)
|
|
|
|
outputs = model.generate(prompts, sampling_params, use_tqdm=False)
|
|
|
|
results = []
|
|
for i, output in enumerate(outputs):
|
|
final_logits = output.outputs[0].logprobs[-1]
|
|
t_logit = final_logits[true_token].logprob if true_token in final_logits else -10.0
|
|
f_logit = final_logits[false_token].logprob if false_token in final_logits else -10.0
|
|
|
|
t_score = math.exp(t_logit)
|
|
f_score = math.exp(f_logit)
|
|
score = t_score / (t_score + f_score)
|
|
results.append({"index": i, "score": score})
|
|
|
|
return {"results": sorted(results, key=lambda x: x['score'], reverse=True)}
|
|
|
|
if __name__ == "__main__":
|
|
import uvicorn
|
|
uvicorn.run(app, host="0.0.0.0", port=8000) |