YS_prediction / model_API_create.py
tianchuang's picture
Update model_API_create.py
bf187e4 verified
Raw
History Blame Contribute Delete
4.27 kB
# 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)