File size: 4,274 Bytes
1361c8e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bf187e4
 
1361c8e
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
# 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)