fiewolf1000/Cross-Encoder
0
1import os2import uuid3import logging4from datetime import datetime5from fastapi import FastAPI, HTTPException, Depends, Request6from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials7from fastapi.responses import HTMLResponse8from pydantic import BaseModel9from transformers import AutoTokenizer, AutoModelForSequenceClassification10import torch11from typing import List, Optional12 13# ------------------- 1. 日志配置 -------------------14logging.basicConfig(15 level=logging.INFO,16 format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",17 datefmt="%Y-%m-%d %H:%M:%S"18)19logger = logging.getLogger("cross-encoder-api")20 21# ------------------- 2. 基础配置(缓存 + 环境变量) -------------------22os.environ["TRANSFORMERS_CACHE"] = "/tmp/huggingface_cache"23os.environ["HUGGINGFACE_HUB_CACHE"] = "/tmp/huggingface_cache"24 25# 从环境变量获取 API Key(OpenAI 风格)26API_KEY = os.getenv("OPENAI_API_KEY")27if not API_KEY:28 logger.error("环境变量 OPENAI_API_KEY 未设置")29 raise ValueError("请设置环境变量 OPENAI_API_KEY")30logger.info("API Key 加载成功")31 32# ------------------- 3. 初始化 FastAPI 应用 -------------------33app = FastAPI(34 title="OpenAI 兼容的 Cross-Encoder 重排序 API",35 description="基于 cross-encoder/ms-marco-MiniLM-L-6-v2 的文本相关性排序接口",36 version="1.0.0"37)38 39# ------------------- 4. OpenAI 风格认证(Bearer Token) -------------------40oauth2_scheme = HTTPBearer(auto_error=False)41 42def verify_api_key(credentials: HTTPAuthorizationCredentials = Depends(oauth2_scheme)):43 """验证 API Key:必须通过 Authorization: Bearer YOUR_API_KEY 传递"""44 request_id = str(uuid.uuid4())[:8] # 生成短请求ID用于日志追踪45 if not credentials:46 logger.warning(f"请求 {request_id}:缺少认证信息")47 raise HTTPException(48 status_code=401,49 detail="缺少认证信息(请使用 'Authorization: Bearer YOUR_API_KEY')",50 headers={"WWW-Authenticate": "Bearer"}51 )52 if credentials.scheme != "Bearer":53 logger.warning(f"请求 {request_id}:认证方案错误,应为 Bearer,实际为 {credentials.scheme}")54 raise HTTPException(55 status_code=401,56 detail="认证方案错误(请使用 'Bearer' 方案)",57 headers={"WWW-Authenticate": "Bearer"}58 )59 if credentials.credentials != API_KEY:60 logger.warning(f"请求 {request_id}:无效的 API Key")61 raise HTTPException(62 status_code=401,63 detail="无效的 API Key",64 headers={"WWW-Authenticate": "Bearer"}65 )66 logger.info(f"请求 {request_id}:API Key 验证通过")67 return (credentials.credentials, request_id) # 返回API Key和请求ID68 69# ------------------- 5. 数据模型定义 -------------------70class RerankRequest(BaseModel):71 query: str72 documents: List[str]73 top_k: Optional[int] = 374 truncation: Optional[bool] = True75 76class DocumentScore(BaseModel):77 document: str78 relevance_score: float79 index: int80 81class RerankResponse(BaseModel):82 request_id: str83 query: str84 top_k: int85 results: List[DocumentScore]86 model: str = "cross-encoder/ms-marco-MiniLM-L-6-v2"87 timestamp: str = datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")[:-3]88 89# GPT 兼容的请求/响应模型90class GPTMessage(BaseModel):91 role: str92 content: str93 94class GPTRequest(BaseModel):95 model: str96 messages: List[GPTMessage]97 top_k: Optional[int] = 398 99class Choice(BaseModel):100 index: int101 message: GPTMessage102 finish_reason: str = "stop"103 104class GPTResponse(BaseModel):105 id: str = f"chatcmpl-{uuid.uuid4().hex}"106 object: str = "chat.completion"107 created: int = int(datetime.now().timestamp())108 model: str109 choices: List[Choice]110 usage: dict = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}111 112# ------------------- 6. 加载 Cross-Encoder 模型 -------------------113class CrossEncoderModel:114 def __init__(self, model_name: str = "cross-encoder/ms-marco-MiniLM-L-6-v2"):115 self.model_name = model_name116 logger.info(f"开始加载模型:{model_name}")117 118 # 验证缓存目录可写119 cache_dir = os.environ.get("TRANSFORMERS_CACHE", "/tmp/huggingface_cache")120 try:121 test_file = os.path.join(cache_dir, "test.txt")122 with open(test_file, "w") as f:123 f.write("test")124 os.remove(test_file)125 logger.info(f"缓存目录可写:{cache_dir}")126 except Exception as e:127 logger.error(f"缓存目录不可写:{str(e)}")128 raise RuntimeError(f"缓存目录不可写:{str(e)}")129 130 # 加载模型131 try:132 logger.info("开始加载分词器...")133 self.tokenizer = AutoTokenizer.from_pretrained(model_name, cache_dir=cache_dir)134 logger.info("分词器加载完成")135 136 logger.info("开始加载模型权重...")137 self.model = AutoModelForSequenceClassification.from_pretrained(model_name, cache_dir=cache_dir)138 logger.info("模型权重加载完成")139 140 self.device = "cuda" if torch.cuda.is_available() else "cpu"141 self.model.to(self.device)142 self.model.eval()143 logger.info(f"模型加载完成,使用设备:{self.device}")144 except Exception as e:145 logger.error(f"模型加载失败:{str(e)}")146 raise147 148 def rerank(self, query: str, documents: List[str], top_k: int, truncation: bool, request_id: str) -> List[DocumentScore]:149 """核心重排序逻辑,增加详细日志"""150 logger.info(f"请求 {request_id}:开始重排序处理,查询长度: {len(query)}, 文档数量: {len(documents)}, top_k: {top_k}")151 152 # 参数校验153 if not documents:154 logger.warning(f"请求 {request_id}:候选文档列表为空")155 raise ValueError("候选文档不能为空")156 if top_k <= 0:157 logger.warning(f"请求 {request_id}:无效的 top_k 值: {top_k}")158 raise ValueError("top_k 必须为正整数")159 160 # 自动将 top_k 限制为文档数量(避免超出)161 adjusted_top_k = min(top_k, len(documents))162 if adjusted_top_k != top_k:163 logger.info(f"请求 {request_id}:top_k 从 {top_k} 调整为 {adjusted_top_k}(文档数量限制)")164 165 # 计算每篇文档的相关性分数166 doc_scores = []167 try:168 for i, doc in enumerate(documents):169 if i % 5 == 0: # 每处理5个文档输出一次日志170 logger.info(f"请求 {request_id}:正在处理第 {i+1}/{len(documents)} 个文档")171 172 inputs = self.tokenizer(173 f"{query} {self.tokenizer.sep_token} {doc}",174 return_tensors="pt",175 padding="max_length",176 truncation=truncation,177 max_length=512178 ).to(self.device)179 180 with torch.no_grad():181 outputs = self.model(**inputs)182 183 score = outputs.logits.item()184 doc_scores.append((doc, score))185 logger.debug(f"请求 {request_id}:文档 {i+1} 分数: {score:.4f}")186 187 # 排序并返回结果188 sorted_docs = sorted(doc_scores, key=lambda x: x[1], reverse=True)[:adjusted_top_k]189 logger.info(f"请求 {request_id}:重排序完成,返回 {len(sorted_docs)} 个结果")190 191 return [192 DocumentScore(document=doc, relevance_score=round(score, 4), index=i)193 for i, (doc, score) in enumerate(sorted_docs)194 ]195 except Exception as e:196 logger.error(f"请求 {request_id}:重排序过程出错: {str(e)}")197 raise198 199# 初始化模型(全局唯一)200try:201 reranker = CrossEncoderModel()202except Exception as e:203 logger.critical(f"模型初始化失败,服务无法启动: {str(e)}")204 raise205 206# ------------------- 7. API 端点(OpenAI 风格路径) -------------------207# 7.1 根路径首页208@app.get("/", response_class=HTMLResponse)209async def home_page(request: Request):210 client_ip = request.client.host211 current_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")212 logger.info(f"首页访问来自 {client_ip}")213 return f"""214<!DOCTYPE html>215<html lang="zh-CN">216<head>217 <meta charset="UTF-8">218 <title>OpenAI 兼容重排序 API</title>219 <style>220 body {{ font-family: Arial, sans-serif; max-width: 1200px; margin: 0 auto; padding: 20px; }}221 h1 {{ color: #2c3e50; border-bottom: 2px solid #3498db; padding-bottom: 10px; }}222 h2 {{ color: #34495e; margin-top: 30px; }}223 pre {{ background: #f8f9fa; padding: 15px; border-radius: 5px; border: 1px solid #e9ecef; overflow-x: auto; }}224 table {{ border-collapse: collapse; width: 100%; margin: 20px 0; }}225 th, td {{ border: 1px solid #e9ecef; padding: 12px; text-align: left; }}226 th {{ background-color: #f1f5f9; }}227 </style>228</head>229<body>230 <h1>OpenAI 兼容的 Cross-Encoder 重排序 API</h1>231 <p>基于 <code>cross-encoder/ms-marco-MiniLM-L-6-v2</code> 模型,支持 OpenAI 风格 API 调用。</p>232 233 <h2>接口列表</h2>234 <table>235 <tr>236 <th>接口</th>237 <th>URL</th>238 <th>方法</th>239 <th>认证</th>240 </tr>241 <tr>242 <td>基础重排序</td>243 <td class="api-url">/v1/rerank</td>244 <td>POST</td>245 <td>Authorization: Bearer API_KEY</td>246 </tr>247 <tr>248 <td>GPT 兼容重排序</td>249 <td class="api-url">/v1/chat/completions</td>250 <td>POST</td>251 <td>Authorization: Bearer API_KEY</td>252 </tr>253 <tr>254 <td>健康检查</td>255 <td class="api-url">/v1/health</td>256 <td>GET</td>257 <td>无需认证</td>258 </tr>259 </table>260 261 <h2>调用示例(Python)</h2>262 <pre><code>import openai263 264client = openai.OpenAI(265 api_key="YOUR_API_KEY",266 base_url="https://your-space.hf.space/v1" # 替换为你的 Space URL267)268 269response = client.chat.completions.create(270 model="cross-encoder/ms-marco-MiniLM-L-6-v2",271 messages=[272 {{273 "role": "user",274 "content": "query: 什么是机器学习?; documents: 机器学习是AI的分支; Python是编程语言;"275 }}276 ],277 top_k=2278)279 280print(response.choices[0].message.content)</code></pre>281</body>282</html>283"""284 285# 7.2 基础重排序接口(/v1/rerank)286@app.post("/v1/rerank", response_model=RerankResponse)287async def base_rerank(288 request: RerankRequest,289 auth_result: tuple = Depends(verify_api_key)290):291 api_key, request_id = auth_result292 try:293 logger.info(f"请求 {request_id}:收到 /v1/rerank 请求,query: {request.query[:50]}...(截断显示)")294 295 # 执行重排序296 results = reranker.rerank(297 query=request.query,298 documents=request.documents,299 top_k=request.top_k,300 truncation=request.truncation,301 request_id=request_id302 )303 304 # 构建响应305 response = RerankResponse(306 request_id=request_id,307 query=request.query,308 top_k=min(request.top_k, len(request.documents)),309 results=results310 )311 312 logger.info(f"请求 {request_id}:处理完成,返回 {len(results)} 个结果")313 return response314 315 except ValueError as e:316 logger.warning(f"请求 {request_id}:参数错误 - {str(e)}")317 raise HTTPException(status_code=400, detail=str(e))318 except Exception as e:319 logger.error(f"请求 {request_id}:服务器错误 - {str(e)}", exc_info=True)320 raise HTTPException(status_code=500, detail=f"服务器错误:{str(e)}")321 322# 7.3 GPT 兼容接口(/v1/chat/completions)323@app.post("/v1/chat/completions", response_model=GPTResponse)324async def gpt_compatible_rerank(325 request: GPTRequest,326 auth_result: tuple = Depends(verify_api_key)327):328 api_key, request_id = auth_result329 try:330 logger.info(f"请求 {request_id}:收到 /v1/chat/completions 请求,模型: {request.model}")331 332 # 验证模型名333 if request.model != reranker.model_name:334 error_msg = f"仅支持模型:{reranker.model_name},实际请求:{request.model}"335 logger.warning(f"请求 {request_id}:{error_msg}")336 raise ValueError(error_msg)337 338 # 验证消息格式339 if not request.messages:340 logger.warning(f"请求 {request_id}:消息列表为空")341 raise ValueError("消息列表不能为空")342 if request.messages[-1].role != "user":343 error_msg = f"最后一条消息必须是 'user' 角色,实际为:{request.messages[-1].role}"344 logger.warning(f"请求 {request_id}:{error_msg}")345 raise ValueError(error_msg)346 347 # 解析输入内容348 content = request.messages[-1].content349 logger.info(f"请求 {request_id}:用户输入: {content[:100]}...(截断显示)")350 351 if "; documents: " not in content:352 error_msg = "输入格式需为 'query: [查询]; documents: [文档1]; [文档2]; ...'"353 logger.warning(f"请求 {request_id}:{error_msg}")354 raise ValueError(error_msg)355 356 query_part, docs_part = content.split("; documents: ")357 query = query_part.replace("query: ", "").strip()358 documents = [doc.strip() for doc in docs_part.split(";") if doc.strip()]359 360 logger.info(f"请求 {request_id}:解析完成,query: {query[:50]}..., 文档数量: {len(documents)}")361 362 # 执行重排序363 results = reranker.rerank(364 query=query,365 documents=documents,366 top_k=request.top_k,367 truncation=True,368 request_id=request_id369 )370 371 # 构建 GPT 风格响应372 response = GPTResponse(373 model=request.model,374 choices=[375 Choice(376 index=0,377 message=GPTMessage(378 role="assistant",379 content=f"重排序结果:{results}"380 )381 )382 ]383 )384 385 logger.info(f"请求 {request_id}:处理完成,返回 {len(results)} 个结果")386 return response387 388 except ValueError as e:389 logger.warning(f"请求 {request_id}:参数错误 - {str(e)}")390 raise HTTPException(status_code=400, detail=str(e))391 except Exception as e:392 logger.error(f"请求 {request_id}:服务器错误 - {str(e)}", exc_info=True)393 raise HTTPException(status_code=500, detail=f"服务器错误:{str(e)}")394 395# 7.4 健康检查接口(/v1/health)396@app.get("/v1/health")397async def health_check(request: Request):398 client_ip = request.client.host399 status = {400 "status": "healthy",401 "model": reranker.model_name,402 "device": reranker.device,403 "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),404 "uptime": datetime.now().strftime("%Y-%m-%d %H:%M:%S") # 简化版uptime405 }406 logger.info(f"健康检查来自 {client_ip}:{status['status']}")407 return status408 409# ------------------- 8. 本地运行入口 -------------------410if __name__ == "__main__":411 import uvicorn412 logger.info("启动本地开发服务器...")413 uvicorn.run(414 app,415 host="0.0.0.0",416 port=7860,417 log_config=None # 使用自定义日志配置418 )419 