""" 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] @app.get("/", response_class=HTMLResponse) def root_ui(): return f""" FinReg BGE Reranker API
Active Microservice

⚖️ FinReg BGE Reranker API

Cloud cross-encoder service powering statutory retrieval re-ranking for Indian Regulatory Compliance.

Active Model: {MODEL_NAME} ({DEVICE.upper()})

Status: {'🟢 Online' if model else '🔴 Loading Error'}

API Endpoints:

POST /rerank (Standard JSON Payload)
GET /health (Healthcheck)
GET /docs (Interactive Swagger API Explorer)
Open Interactive API Docs (Swagger) →
""" @app.get("/health") def health_check(): return { "status": "healthy" if model else "unhealthy", "model": MODEL_NAME, "device": DEVICE, "model_loaded": model is not None } @app.post("/rerank", response_model=RerankResponse) 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)