Spaces:
Sleeping
Sleeping
| import time | |
| from fastapi import APIRouter, HTTPException | |
| from typing import Dict, Any | |
| from app.schemas.prediction import ( | |
| PredictRequest, | |
| PredictResponse, | |
| OverallPrediction, | |
| TimingInfo, | |
| PredictMetadata, | |
| ModelVersions | |
| ) | |
| from app.services.model_loader import model_manager | |
| from app.services.predictor import predict_lr, predict_lstm, predict_bert | |
| router = APIRouter() | |
| MODEL_VERSIONS = ModelVersions( | |
| lr="1.0.0", | |
| lstm="1.0.0", | |
| bert="1.0.0" | |
| ) | |
| async def predict(request: PredictRequest): | |
| start_time = time.perf_counter() | |
| models_result = {} | |
| pos_count = 0 | |
| neg_count = 0 | |
| total_count = 0 | |
| errors = [] | |
| def try_predict(model_key: str, predict_fn, text: str): | |
| nonlocal pos_count, neg_count, total_count | |
| try: | |
| result = predict_fn(text) | |
| if result is None: | |
| errors.append(f"{model_key}: returned None") | |
| return | |
| models_result[model_key] = result | |
| total_count += 1 | |
| if result.label == "Positive": | |
| pos_count += 1 | |
| else: | |
| neg_count += 1 | |
| except Exception as e: | |
| import traceback | |
| errors.append(f"{model_key}: {str(e)}\n{traceback.format_exc()}") | |
| if "lr" in request.models: | |
| try_predict("lr", predict_lr, request.text) | |
| if "lstm" in request.models: | |
| try_predict("lstm", predict_lstm, request.text) | |
| if "bert" in request.models: | |
| try_predict("bert", predict_bert, request.text) | |
| if total_count == 0: | |
| raise HTTPException(status_code=503, detail=f"No models available: {'; '.join(errors)}") | |
| majority_label = "Positive" if pos_count >= neg_count else "Negative" | |
| agreement = f"{max(pos_count, neg_count)}/{total_count}" | |
| total_ms = (time.perf_counter() - start_time) * 1000 | |
| return PredictResponse( | |
| overall=OverallPrediction(label=majority_label, agreement=agreement), | |
| models=models_result, | |
| timing=TimingInfo(total_ms=total_ms), | |
| metadata=PredictMetadata(model_versions=MODEL_VERSIONS) | |
| ) if not errors else { | |
| "overall": OverallPrediction(label=majority_label, agreement=agreement).model_dump(), | |
| "models": {k: v.model_dump() for k, v in models_result.items()}, | |
| "timing": TimingInfo(total_ms=total_ms).model_dump(), | |
| "metadata": PredictMetadata(model_versions=MODEL_VERSIONS).model_dump(), | |
| "errors": errors, | |
| } | |
| async def health_check() -> Dict[str, Any]: | |
| return { | |
| "status": "healthy", | |
| "backend": False, | |
| "ml_service": True, | |
| "database": False, | |
| "models_loaded": { | |
| "logistic_regression": model_manager.models_loaded.get("logistic_regression", False), | |
| "lstm": model_manager.models_loaded.get("lstm", False), | |
| "bert": model_manager.models_loaded.get("bert", False) | |
| }, | |
| "load_errors": model_manager.load_errors, | |
| "version": "1.0.0" | |
| } | |
| async def get_models() -> Dict[str, Any]: | |
| return { | |
| "lr": { | |
| "name": "Logistic Regression (TF-IDF)", | |
| "type": "Machine Learning", | |
| "version": MODEL_VERSIONS.lr, | |
| "description": "Bag-of-words classifier using TF-IDF features" | |
| }, | |
| "lstm": { | |
| "name": "Bi-LSTM", | |
| "type": "Deep Learning", | |
| "version": MODEL_VERSIONS.lstm, | |
| "description": "Bidirectional recurrent neural network" | |
| }, | |
| "bert": { | |
| "name": "BERT (Fine-Tuned)", | |
| "type": "Transformer", | |
| "version": MODEL_VERSIONS.bert, | |
| "description": "Fine-tuned contextual embedding model" | |
| } | |
| } | |
| async def get_model_metrics() -> Dict[str, Any]: | |
| return { | |
| # Minimal mock payload for now, database will hold the real ones | |
| "status": "Not implemented here, served by DB via Gateway" | |
| } | |