"""
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)