Spaces:
Sleeping
Sleeping
| # 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") | |
| 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) |