Nkosiii's picture
Upload 44 files
5863f1d verified
Raw
History Blame Contribute Delete
4.63 kB
# src/inference.py
import joblib
import pickle
from pathlib import Path
from typing import Any, List, Optional
import numpy as np
import pandas as pd
import logging
logger = logging.getLogger(__name__)
DEFAULT_MODEL_PATH = Path("data/electricityDemandModel.pkl")
DEFAULT_FEATURE_PATH = Path("data/featureColumns.pkl")
class ModelPredictor:
"""
Robust model wrapper to load saved model + feature list and produce predictions.
Methods:
- __init__(model_path, feature_path)
- predict(X) -> np.ndarray
- explain(X) -> shap values | None
"""
def __init__(self, model_path: str = None, feature_path: str = None):
self.model_path = Path(model_path) if model_path else DEFAULT_MODEL_PATH
self.feature_path = Path(feature_path) if feature_path else DEFAULT_FEATURE_PATH
self.model: Optional[Any] = None
self.feature_columns: Optional[List[str]] = None
self._load_model()
self._load_features()
def _load_model(self):
if not self.model_path.exists():
raise FileNotFoundError(f"Model file not found at {self.model_path}")
# try joblib first
try:
self.model = joblib.load(self.model_path)
return
except Exception:
logger.debug("joblib.load failed; trying pickle", exc_info=True)
# try pickle
try:
with open(self.model_path, "rb") as f:
self.model = pickle.load(f)
return
except Exception as e:
logger.exception("Failed to load model", exc_info=True)
raise RuntimeError(f"Failed to load model from {self.model_path}: {e}")
def _load_features(self):
if not self.feature_path.exists():
# feature list optional; downstream code must handle None
self.feature_columns = None
return
try:
obj = joblib.load(self.feature_path)
except Exception:
try:
with open(self.feature_path, "rb") as f:
obj = pickle.load(f)
except Exception as e:
logger.exception("Failed to load feature columns", exc_info=True)
raise RuntimeError(f"Failed to load feature columns: {e}")
if isinstance(obj, (list, tuple)):
self.feature_columns = list(obj)
elif isinstance(obj, (pd.Series, np.ndarray)):
self.feature_columns = list(obj)
else:
raise RuntimeError("featureColumns.pkl must contain a list-like object")
def predict(self, X: pd.DataFrame) -> np.ndarray:
"""
Predict using the loaded model. X should be a pandas DataFrame aligned with feature_columns.
Returns a 1D numpy array.
"""
if self.model is None:
raise RuntimeError("Model not loaded")
# If feature_columns provided, align X
if self.feature_columns is not None:
missing = [c for c in self.feature_columns if c not in X.columns]
if missing:
# create missing cols with zeros
for c in missing:
X[c] = 0
X = X[self.feature_columns]
try:
preds = self.model.predict(X)
return np.ravel(preds)
except Exception as e:
logger.exception("Model prediction failed", exc_info=True)
# final fallback: try numpy interface
try:
return np.ravel(self.model.predict(np.asarray(X)))
except Exception as e2:
raise RuntimeError(f"Prediction failed: {e} ; fallback also failed: {e2}")
def explain(self, X):
"""
Attempt to compute SHAP explanations. Return None if shap not installed or fails.
"""
try:
import shap
except Exception:
logger.warning("shap not installed; explain() will return None")
return None
try:
# try common explainer patterns
explainer = None
try:
explainer = shap.Explainer(self.model, X)
shap_values = explainer(X)
return shap_values
except Exception:
# fallback to TreeExplainer
explainer = shap.TreeExplainer(self.model)
return explainer.shap_values(X)
except Exception:
logger.exception("SHAP explanation failed", exc_info=True)
return None