MatteoAldovardi
a
17a9bed
Raw
History Blame Contribute Delete
3.23 kB
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
import os
import joblib
import pandas as pd
import threading
import time
import requests
from huggingface_hub import hf_hub_download
# --- 0. Initialize FastAPI app ---
app = FastAPI()
@app.get("/")
def read_root():
return {
"message": "Welcome to the Taxi Fare Prediction API. Use POST /predict to get predictions."
}
@app.get("/health")
def health():
return {"status": "ok"}
# --- 1. Define Input and Output Data Models ---
class InferenceInput(BaseModel):
vendor_id: str
dist_meters: float
wait_sec: float
geodetic_dist: float
mean_velocity: float
is_rush_hour: bool
model_name: str # "bog", "mex", or "uio"
class InferenceOutput(BaseModel):
trip_duration: float
model_used: str
message: str
# --- 2. Lazy-loading ML Model Manager ---
class MLModels:
def __init__(self):
self.models = {}
self.repo_id = "MatteoAld/my-ml-models" # ⚠️ Change to your actual repo
self.valid_models = {"bog", "mex", "uio"}
self.file_template = "{}_ridge_pipeline.pkl"
def get_model(self, model_name: str):
model_name = model_name.lower()
if model_name not in self.valid_models:
raise ValueError(
f"Invalid model name '{model_name}'. Choose from 'bog', 'mex', or 'uio'."
)
if model_name in self.models:
return self.models[model_name]
filename = self.file_template.format(model_name)
try:
# Download from Hugging Face Hub (cached automatically)
model_path = hf_hub_download(
repo_id=self.repo_id,
filename=filename,
cache_dir="/tmp/hf_models" # Use /tmp for write permissions
)
model = joblib.load(model_path)
self.models[model_name] = model
print(f"✅ Downloaded and loaded model '{model_name}' from Hugging Face Hub")
return model
except Exception as e:
raise RuntimeError(f"Error loading model '{model_name}': {e}")
def predict_one(self, features: dict, model_name: str):
model = self.get_model(model_name)
X_df = pd.DataFrame([features])
pred = model.predict(X_df)[0]
return float(pred)
# --- 3. Instantiate ML Model Manager ---
ml_models = MLModels()
# --- 4. Inference Endpoint ---
@app.post("/predict", response_model=InferenceOutput)
async def predict_inference(data: InferenceInput):
try:
features = data.dict()
model_name = features.pop("model_name")
trip_duration = ml_models.predict_one(features, model_name)
return InferenceOutput(
trip_duration=trip_duration,
model_used=model_name,
message=f"Inference successful using {model_name.upper()} model.",
)
except FileNotFoundError as fnf:
raise HTTPException(status_code=404, detail=str(fnf))
except ValueError as ve:
raise HTTPException(status_code=400, detail=str(ve))
except Exception as e:
raise HTTPException(status_code=500, detail=f"Inference failed: {str(e)}")
# --- 5. Run the app ---