Download app.py from Mister2005/Cross-Encoder-Reranking-API: direct link, hf CLI and curl.
- Browser
- Download file 6.12 kB
-
https://huggingface.co/spaces/Mister2005/Cross-Encoder-Reranking-API/resolve/main/app.py
- Command line
-
hf download hf://spaces/Mister2005/Cross-Encoder-Reranking-API/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Mister2005/Cross-Encoder-Reranking-API/resolve/main/app.py
6.12 kB
| """ | |
| FinReg BGE Cross-Encoder Reranking API. | |
| Hosted on Hugging Face Spaces (Mister2005/Cross-Encoder-Reranking-API). | |
| Default Model: BAAI/bge-reranker-base (State-of-the-art Chinese & English Cross-Encoder) | |
| Framework: FastAPI + SentenceTransformers / PyTorch CPU/GPU | |
| """ | |
| import os | |
| import time | |
| import torch | |
| import uvicorn | |
| from typing import List, Dict, Any, Optional | |
| from fastapi import FastAPI, HTTPException | |
| from fastapi.responses import HTMLResponse | |
| from pydantic import BaseModel, Field | |
| from sentence_transformers import CrossEncoder | |
| MODEL_NAME = os.getenv("MODEL_NAME", "BAAI/bge-reranker-base") | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| app = FastAPI( | |
| title="FinReg BGE Reranker API", | |
| description="High-Precision Regulatory Document Re-Ranking powered by BAAI/bge-reranker-base.", | |
| version="2.0.0" | |
| ) | |
| print(f"Loading CrossEncoder model '{MODEL_NAME}' on {DEVICE}...") | |
| try: | |
| model = CrossEncoder(MODEL_NAME, max_length=512, device=DEVICE) | |
| print(f"Successfully initialized {MODEL_NAME}!") | |
| except Exception as e: | |
| print(f"Error loading model: {e}") | |
| model = None | |
| class RerankItem(BaseModel): | |
| rank: int | |
| original_index: int | |
| score: float | |
| document: str | |
| class RerankRequest(BaseModel): | |
| query: str = Field(..., description="The search query or compliance question") | |
| documents: List[str] = Field(..., description="List of candidate text passages to re-rank") | |
| top_k: Optional[int] = Field(default=None, description="Number of top passages to return (defaults to all)") | |
| class RerankResponse(BaseModel): | |
| query: str | |
| model: str | |
| total_evaluated: int | |
| latency_ms: float | |
| scores: List[float] | |
| ranked_indices: List[int] | |
| ranked_results: List[RerankItem] | |
| def root_ui(): | |
| return f""" | |
| <!DOCTYPE html> | |
| <html> | |
| <head> | |
| <title>FinReg BGE Reranker API</title> | |
| <meta charset="utf-8"> | |
| <style> | |
| body {{ font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; max-width: 800px; margin: 40px auto; padding: 0 20px; color: #1e293b; background: #f8fafc; }} | |
| .card {{ background: white; padding: 30px; border-radius: 12px; box-shadow: 0 4px 6px -1px rgb(0 0 0 / 0.1); }} | |
| h1 {{ color: #0f172a; margin-top: 0; }} | |
| .badge {{ display: inline-block; padding: 4px 10px; border-radius: 9999px; font-size: 12px; font-weight: 600; background: #e0e7ff; color: #3730a3; }} | |
| .endpoint {{ background: #f1f5f9; padding: 12px; border-radius: 6px; font-family: monospace; margin: 12px 0; }} | |
| a.btn {{ display: inline-block; background: #2563eb; color: white; padding: 10px 18px; border-radius: 6px; text-decoration: none; font-weight: 500; margin-top: 15px; }} | |
| a.btn:hover {{ background: #1d4ed8; }} | |
| </style> | |
| </head> | |
| <body> | |
| <div class="card"> | |
| <span class="badge">Active Microservice</span> | |
| <h1>⚖️ FinReg BGE Reranker API</h1> | |
| <p>Cloud cross-encoder service powering statutory retrieval re-ranking for Indian Regulatory Compliance.</p> | |
| <p><strong>Active Model:</strong> <code>{MODEL_NAME}</code> ({DEVICE.upper()})</p> | |
| <p><strong>Status:</strong> {'🟢 Online' if model else '🔴 Loading Error'}</p> | |
| <h3>API Endpoints:</h3> | |
| <div class="endpoint">POST /rerank (Standard JSON Payload)</div> | |
| <div class="endpoint">GET /health (Healthcheck)</div> | |
| <div class="endpoint">GET /docs (Interactive Swagger API Explorer)</div> | |
| <a class="btn" href="/docs">Open Interactive API Docs (Swagger) →</a> | |
| </div> | |
| </body> | |
| </html> | |
| """ | |
| def health_check(): | |
| return { | |
| "status": "healthy" if model else "unhealthy", | |
| "model": MODEL_NAME, | |
| "device": DEVICE, | |
| "model_loaded": model is not None | |
| } | |
| def rerank_documents(request: RerankRequest): | |
| if not model: | |
| raise HTTPException(status_code=503, detail="Model is not loaded on server.") | |
| if not request.query.strip() or not request.documents: | |
| return RerankResponse( | |
| query=request.query, | |
| model=MODEL_NAME, | |
| total_evaluated=0, | |
| latency_ms=0.0, | |
| scores=[], | |
| ranked_indices=[], | |
| ranked_results=[] | |
| ) | |
| start_time = time.time() | |
| try: | |
| # Create (query, doc) pairs | |
| pairs = [[request.query, doc] for doc in request.documents] | |
| # CrossEncoder scoring with sigmoid activation | |
| raw_scores = model.predict(pairs, convert_to_numpy=True, show_progress_bar=False) | |
| scores_list = [float(s) for s in raw_scores] | |
| # Sigmoid normalization: 1 / (1 + exp(-score)) | |
| probs = [round(float(torch.sigmoid(torch.tensor(s)).item()), 4) for s in scores_list] | |
| # Rank pairs | |
| indexed = list(enumerate(probs)) | |
| indexed.sort(key=lambda x: x[1], reverse=True) | |
| ranked_indices = [idx for idx, _ in indexed] | |
| top_k = request.top_k if request.top_k and request.top_k > 0 else len(request.documents) | |
| ranked_results = [] | |
| for rank_num, (orig_idx, score) in enumerate(indexed[:top_k], 1): | |
| ranked_results.append(RerankItem( | |
| rank=rank_num, | |
| original_index=orig_idx, | |
| score=score, | |
| document=request.documents[orig_idx] | |
| )) | |
| latency = (time.time() - start_time) * 1000.0 | |
| return RerankResponse( | |
| query=request.query, | |
| model=MODEL_NAME, | |
| total_evaluated=len(request.documents), | |
| latency_ms=round(latency, 2), | |
| scores=probs, | |
| ranked_indices=ranked_indices, | |
| ranked_results=ranked_results | |
| ) | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| if __name__ == "__main__": | |
| uvicorn.run(app, host="0.0.0.0", port=7860) | |