Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7,042 changes: 5,949 additions & 1,093 deletions pixi.lock

Large diffs are not rendered by default.

8 changes: 8 additions & 0 deletions pixi.toml
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,9 @@ test-server = "python run_fastapi.py"
test = "python tests/test_fastapi.py"
docs = "echo 'API docs available at http://localhost:8888/docs'"

[environments]
test = { features = ["test"] }

[activation.env]
HOST="0.0.0.0"
PORT="8888"
Expand All @@ -23,6 +26,7 @@ REDIS_PORT="6379"
python = "3.12.*"
redis-py = ">=5.0.1,<6"
requests = ">=2.31.0,<3"
pytest = ">=9.0.2,<10"

[pypi-dependencies]
fastapi = ">=0.104.1,<0.105"
Expand All @@ -33,6 +37,10 @@ slowapi = ">=0.1.9,<0.2"
umap-learn = ">=0.5.5,<0.6"
typing_extensions = ">=4.14.1,<5"
jarvais = ">=0.20.0, <0.21"
httpx = ">=0.24,<0.28"

[feature.test.pypi-dependencies]
pytest = "*"



Expand Down
5 changes: 3 additions & 2 deletions src/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from .config import settings
from .storage import storage_manager
from .routers import upload, visualization, analyzers, trainers, health, dashboard
from .routers import upload, visualization, analyzers, trainers, health, dashboards, explainers

# Configure logging
logging.basicConfig(
Expand Down Expand Up @@ -87,7 +87,8 @@ def decorator(func):
app.include_router(visualization.router)
app.include_router(analyzers.router)
app.include_router(trainers.router)
app.include_router(dashboard.router)
app.include_router(dashboards.router)
app.include_router(explainers.router)
app.include_router(health.router)


Expand Down
7 changes: 6 additions & 1 deletion src/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,4 +102,9 @@ class InferenceResult(BaseModel):
predictions: List[Any] = Field(..., description="Predicted values")
probabilities: Optional[List[List[float]]] = Field(None, description="Prediction probabilities for classification tasks")
num_samples: int = Field(..., description="Number of samples predicted")
created_at: str = Field(..., description="Creation timestamp")
created_at: str = Field(..., description="Creation timestamp")


#
# NOTE: Highcharts response wrapper removed.
# Endpoints now return raw chart payloads directly.
6 changes: 3 additions & 3 deletions src/modules/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,9 @@
This package contains helper functions and utilities for training and inference.
"""

from . import jarvais_train
from . import jarvais_infer
from . import train
from . import infer
from . import plot

__all__ = ['jarvais_train', 'jarvais_infer', 'plot']
__all__ = ['train', 'infer', 'plot']

161 changes: 161 additions & 0 deletions src/modules/explain.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
import logging
from typing import Any, Tuple

import numpy as np
import pandas as pd
from fastapi import HTTPException

from ..storage import storage_manager

logger = logging.getLogger(__name__)


def _encode_frame_for_shap(df: pd.DataFrame) -> Tuple[pd.DataFrame, dict[str, pd.Index]]:
"""Convert non-numeric columns to categorical codes for SHAP."""
encoded = df.copy()
cat_maps: dict[str, pd.Index] = {}
for col in encoded.columns:
if not pd.api.types.is_numeric_dtype(encoded[col]):
cat = pd.Categorical(encoded[col])
cat_maps[col] = cat.categories
encoded[col] = cat.codes.astype(float)
return encoded, cat_maps


def _decode_frame_from_shap(df: pd.DataFrame, cat_maps: dict[str, pd.Index]) -> pd.DataFrame:
"""Restore categorical codes back to original labels after SHAP perturbations."""
decoded = df.copy()
for col, categories in cat_maps.items():
# SHAP returns floats; round to nearest int code and map back
decoded[col] = (
decoded[col]
.round()
.astype(int)
.map(lambda idx: categories[idx] if 0 <= idx < len(categories) else None)
)
return decoded


def load_trainer(trainer_id: str):
"""Load a TrainerSupervised instance from stored metadata."""
if not storage_manager.check_trainer(trainer_id):
raise HTTPException(status_code=404, detail="Trainer not found")

trainer_data = storage_manager.get_trainer(trainer_id)
if not trainer_data:
raise HTTPException(status_code=404, detail="Trainer not found")

output_dir = trainer_data.get("output_dir")
if not output_dir:
raise HTTPException(status_code=500, detail="Trainer output directory not set")

try:
from jarvais.trainer import TrainerSupervised

return TrainerSupervised.load_trainer(output_dir)
except FileNotFoundError as exc:
logger.error("Trainer files missing at %s: %s", output_dir, exc)
raise HTTPException(status_code=500, detail="Trainer artifacts missing on disk") from exc
except Exception as exc: # pragma: no cover
logger.error("Failed to load trainer %s: %s", trainer_id, exc)
raise HTTPException(status_code=500, detail="Failed to load trainer") from exc


def sample_frame(df: pd.DataFrame, max_rows: int = 200) -> pd.DataFrame:
"""Return a small, reproducible sample for expensive computations."""
if df.empty:
raise HTTPException(status_code=400, detail="Trainer data is empty")
if len(df) <= max_rows:
return df
return df.sample(n=max_rows, random_state=0)


def select_feature_frame(trainer) -> Tuple[pd.DataFrame, pd.DataFrame]:
"""Return X (features) and y (target) for downstream explainer utilities."""
X = getattr(trainer, "X_train", None)
y = getattr(trainer, "y_train", None)
if X is None or y is None:
raise HTTPException(status_code=500, detail="Trainer is missing training data")
return X, y


def predict_mean(trainer, frame: pd.DataFrame) -> float:
"""Average prediction helper across tasks."""
predictor = trainer.predictor
if getattr(predictor, "can_predict_proba", False) and trainer.settings.task in {"binary", "multiclass"}:
proba = predictor.predict_proba(frame)
target_col = proba.columns[1] if proba.shape[1] > 1 else proba.columns[0]
return float(proba[target_col].mean())
preds = predictor.predict(frame, as_pandas=False)
return float(np.asarray(preds).mean())


def compute_shap(trainer, sample_size: int = 200) -> Tuple[Any, pd.DataFrame]:
"""Compute SHAP values on a small sample."""
try:
import shap # type: ignore
except ImportError as exc: # pragma: no cover
raise HTTPException(status_code=500, detail="SHAP is not available in the environment") from exc

X, _ = select_feature_frame(trainer)
X_sample = sample_frame(X, max_rows=sample_size)

# SHAP's tabular masker expects numeric input; encode categoricals and
# decode them back before passing to the predictor.
X_encoded, cat_maps = _encode_frame_for_shap(X_sample)
cols = X_encoded.columns
predictor = trainer.predictor
is_clf = getattr(predictor, "can_predict_proba", False) and trainer.settings.task in {"binary", "multiclass"}

def predict_fn(data):
encoded_df = pd.DataFrame(data, columns=cols)
decoded_df = _decode_frame_from_shap(encoded_df, cat_maps)
if is_clf:
return predictor.predict_proba(decoded_df).values
return predictor.predict(decoded_df, as_pandas=False)

explainer = shap.Explainer(predict_fn, X_encoded)
shap_values = explainer(X_encoded)

# Return the original sample (with human-friendly columns) for charting
return shap_values, X_sample


def compute_pdp(trainer, feature: str, grid_size: int = 15) -> Tuple[np.ndarray, np.ndarray]:
"""Compute a simple 1D partial dependence using model predictions."""
X, _ = select_feature_frame(trainer)
if feature not in X.columns:
raise HTTPException(status_code=400, detail=f"Feature '{feature}' not found")

base = sample_frame(X, max_rows=300).copy()
series = base[feature]
if pd.api.types.is_numeric_dtype(series):
grid = np.linspace(series.min(), series.max(), grid_size)
else:
grid = series.dropna().unique()

preds = []
for value in grid:
modified = base.copy()
modified[feature] = value
preds.append(predict_mean(trainer, modified))

return np.asarray(grid), np.asarray(preds)


def compute_permutation_importance(trainer) -> pd.DataFrame:
"""Proxy to AutoGluon feature_importance with a safe default."""
predictor = trainer.predictor
try:
X, y = select_feature_frame(trainer)
frame = pd.concat([X, y], axis=1)
importance = predictor.feature_importance(frame)
except Exception as exc:
logger.error("Permutation importance failed: %s", exc)
raise HTTPException(status_code=500, detail="Failed to compute permutation importance") from exc

if not isinstance(importance, pd.DataFrame):
raise HTTPException(status_code=500, detail="Unexpected importance format from predictor")
if "importance" not in importance.columns:
raise HTTPException(status_code=500, detail="Predictor importance response missing 'importance' column")
return importance
File renamed without changes.
16 changes: 15 additions & 1 deletion src/modules/plot/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,14 @@
from .umap import get_umap_json
from .violinplot import get_violin_plot_json
from .boxplot import get_box_plot_json, get_grouped_box_plot_json
from .dashboard import get_dashboard_json
from .analyzer_dashboard import get_dashboard_json
from .distribution import get_distribution_chart
from .explainer_dashboard import get_explainer_dashboard
from .pdp import get_partial_dependence_chart
from .permutation_importance import get_permutation_importance_chart
from .shap_bar import get_shap_bar_chart
from .shap_dependence import get_shap_dependence_chart
from .shap_summary import get_shap_summary_chart

__all__ = [
"get_corr_heatmap_json",
Expand All @@ -15,4 +22,11 @@
"get_box_plot_json",
"get_grouped_box_plot_json",
"get_dashboard_json",
"get_explainer_dashboard",
"get_distribution_chart",
"get_partial_dependence_chart",
"get_permutation_importance_chart",
"get_shap_bar_chart",
"get_shap_dependence_chart",
"get_shap_summary_chart",
]
File renamed without changes.
45 changes: 45 additions & 0 deletions src/modules/plot/distribution.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
from __future__ import annotations

from typing import Any, Dict

import numpy as np
import pandas as pd


def get_distribution_chart(
train_series: pd.Series,
test_series: pd.Series,
feature_name: str,
bins: int = 15,
) -> Dict[str, Any]:
"""Return overlapping histograms (numeric) or bar chart (categorical) for train vs test."""
if pd.api.types.is_numeric_dtype(train_series):
combined = pd.concat([train_series, test_series])
_, bin_edges = np.histogram(combined.dropna(), bins=bins)
train_hist, _ = np.histogram(train_series.dropna(), bins=bin_edges)
test_hist, _ = np.histogram(test_series.dropna(), bins=bin_edges)
categories = [f"{bin_edges[i]:.2f}-{bin_edges[i+1]:.2f}" for i in range(len(bin_edges) - 1)]
series = [
{"name": "train", "data": train_hist.tolist(), "type": "column", "pointPadding": 0},
{"name": "test", "data": test_hist.tolist(), "type": "column", "pointPadding": 0.1},
]
else:
train_counts = train_series.value_counts()
test_counts = test_series.value_counts()
categories = sorted(set(train_counts.index).union(set(test_counts.index)))
series = [
{"name": "train", "data": [int(train_counts.get(cat, 0)) for cat in categories]},
{"name": "test", "data": [int(test_counts.get(cat, 0)) for cat in categories]},
]

config = {
"chart": {"type": "column"},
"title": {"text": f"Train vs Test Distribution: {feature_name}"},
"xAxis": {"categories": categories, "title": {"text": feature_name}},
"yAxis": {"title": {"text": "Count"}},
"plotOptions": {"column": {"grouping": True, "shadow": False}},
"series": series,
}

return config

65 changes: 65 additions & 0 deletions src/modules/plot/explainer_dashboard.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
from __future__ import annotations

from typing import Any, Dict, List

import numpy as np

from .distribution import get_distribution_chart
from .pdp import get_partial_dependence_chart
from .permutation_importance import get_permutation_importance_chart
from .shap_bar import get_shap_bar_chart
from .shap_dependence import get_shap_dependence_chart
from .shap_summary import get_shap_summary_chart
from ...utils.shap import to_numpy, top_indices_by_mean_abs


def get_explainer_dashboard(
trainer,
sample_size: int = 200,
top_n: int = 10,
grid_size: int = 15,
) -> List[Dict[str, Any]]:
"""
Generate a lightweight explainer dashboard as a list of Highcharts configs.
The dashboard includes:
- SHAP summary (beeswarm)
- SHAP bar importance
- SHAP dependence for the top feature
- Partial dependence for the top feature
- Permutation importance
- Train/Test distribution for the top feature
"""
charts: List[Dict[str, Any]] = []

shap_values, X_sample = compute_shap(trainer, sample_size=sample_size)
shap_matrix = to_numpy(shap_values)
feature_names = list(X_sample.columns)
top_idx = top_indices_by_mean_abs(shap_matrix, top_n)
top_feature = feature_names[top_idx[0]] if top_idx else feature_names[0]

charts.append(get_shap_summary_chart(shap_values, feature_names, top_n=top_n))
charts.append(get_shap_bar_chart(shap_values, feature_names, top_n=top_n))

shap_feat_values = shap_matrix[:, feature_names.index(top_feature)]
charts.append(
get_shap_dependence_chart(
X_sample[top_feature],
shap_feat_values,
top_feature,
)
)

grid, preds = compute_pdp(trainer, top_feature, grid_size=grid_size)
charts.append(get_partial_dependence_chart(grid, preds, feature_name=top_feature, task=trainer.settings.task))

importance_df = compute_permutation_importance(trainer)
charts.append(get_permutation_importance_chart(importance_df, top_n=top_n))

X_train = getattr(trainer, "X_train", None)
X_test = getattr(trainer, "X_test", None)
if X_train is not None and X_test is not None and top_feature in X_train.columns and top_feature in X_test.columns:
charts.append(get_distribution_chart(X_train[top_feature], X_test[top_feature], feature_name=top_feature))

return charts


Loading
Loading