# app.py from fastapi import FastAPI, HTTPException from pydantic import BaseModel from typing import Optional import torch import numpy as np from joblib import load import os # Import your existing components from model_API import BertEmbedder, FeatureExtractor, MCDropoutModel, mc_dropout_predict # ---------------------------- # Configuration # ---------------------------- DEVICE = "cuda" if torch.cuda.is_available() else "cpu" # Model paths (adjust if needed) SCALER_PATH = "./ANN-modeling/mc_dropout_scaler.pkl" MODEL_PATH = './ANN-modeling/mc_dropout_best_vam_only.pth' EMBEDDING_MODEL_PATH = "./embedding-model/matscibert" FEATURE_EXTRACTOR_PATH = "./DANN-modeling/feature_extractor.pth" # Validate required files exist for path in [SCALER_PATH, MODEL_PATH, EMBEDDING_MODEL_PATH, FEATURE_EXTRACTOR_PATH]: if not os.path.exists(path): raise FileNotFoundError(f"Required model file not found: {path}") # ---------------------------- # Load models once at startup # ---------------------------- print("Loading models...") try: embedder = BertEmbedder(model_name=EMBEDDING_MODEL_PATH, device=DEVICE) extractor = FeatureExtractor().to(DEVICE) extractor.load_state_dict(torch.load(FEATURE_EXTRACTOR_PATH, map_location=DEVICE)) extractor.eval() predictor = MCDropoutModel().to(DEVICE) predictor.load_state_dict(torch.load(MODEL_PATH, map_location=DEVICE)) predictor.eval() scaler = load(SCALER_PATH) print("All models loaded successfully.") except Exception as e: raise RuntimeError(f"Failed to initialize models: {e}") # ---------------------------- # API Schema # ---------------------------- class PredictionRequest(BaseModel): Al_content: Optional[float] = None Nb_content: Optional[float] = None Ta_content: Optional[float] = None Ti_content: Optional[float] = None Zr_content: Optional[float] = None Mo_content: Optional[float] = None V_content: Optional[float] = None Cr_content: Optional[float] = None W_content: Optional[float] = None Hf_content: Optional[float] = None Ni_content: Optional[float] = None process_text: Optional[str] = None class PredictionResponse(BaseModel): prediction: float prediction_std: float # ---------------------------- # FastAPI App # ---------------------------- app = FastAPI(title="AMed-LWRHEAs Property Prediction API", version="1.0") @app.post("/predict", response_model=PredictionResponse) def predict_api(request: PredictionRequest): try: # Convert None → 0.0 for all elements comp_list = [ request.Al_content or 0.0, request.Nb_content or 0.0, request.Ta_content or 0.0, request.Ti_content or 0.0, request.Zr_content or 0.0, request.Mo_content or 0.0, request.V_content or 0.0, request.Cr_content or 0.0, request.W_content or 0.0, request.Hf_content or 0.0, request.Ni_content or 0.0, ] comp_list = [float(x) for x in comp_list] # Handle process embedding if request.process_text is not None: emb = embedder.encode(request.process_text) proc_emb = np.array(emb) # Combine features comp_array = np.atleast_2d(comp_list) # (1, 11) proc_emb_array = np.tile(proc_emb, (comp_array.shape[0], 1)) # (1, 768) X = np.hstack([comp_array, proc_emb_array]) # (1, 779) # Extract latent features with torch.no_grad(): latent = extractor(torch.FloatTensor(X).to(DEVICE)).cpu().numpy() # Scale and predict with MC Dropout latent_scaled = scaler.transform(latent) pred_ys, pred_ys_std = mc_dropout_predict(predictor, latent_scaled) return PredictionResponse( prediction=np.round(float(pred_ys[0]), 2), prediction_std=np.round(float(pred_ys_std[0]), 2) ) except Exception as e: raise HTTPException(status_code=500, detail=f"Prediction error: {str(e)}") if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=8000)