Team Ai
Apppublic

fiewolf1000/Cross-Encoder

sourceHugging Faceotherupdated 1y agoView on Hugging Face
0likes
app.py419 linesDownload Raw Back to root
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