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
2 changes: 0 additions & 2 deletions DashAI/back/api/api_v1/endpoints/predict.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,7 +252,6 @@ async def delete_prediction(
@inject
async def preview_manual_prediction(
request: Request,
component_registry: "ComponentRegistry" = Depends(lambda: di["component_registry"]),
session_factory: "sessionmaker" = Depends(lambda: di["session_factory"]),
):
"""Run a synchronous manual prediction and return results without persisting.
Expand Down Expand Up @@ -335,7 +334,6 @@ async def preview_manual_prediction(
run_manual_prediction,
run_id=run_id_int,
manual_input_data=rows_data,
component_registry=component_registry,
session_factory=session_factory,
)
return {"columns": columns, "rows": rows}
7 changes: 5 additions & 2 deletions DashAI/back/explainability/explainers/contrastive_shap.py
Original file line number Diff line number Diff line change
Expand Up @@ -238,7 +238,10 @@ def fit(
"""
import shap

from DashAI.back.explainability.model_input import prepare_model_input
from DashAI.back.explainability.model_input import (
as_shap_predictor,
prepare_model_input,
)

x, y = background_dataset
# SHAP calls the model with perturbed frames, which skip the model
Expand All @@ -254,7 +257,7 @@ def fit(
background_data = shap.sample(background_data, n_samples)

self.explainer = shap.KernelExplainer(
model=self.model.predict,
model=as_shap_predictor(self.model),
data=background_data,
feature_names=feature_names,
)
Expand Down
7 changes: 5 additions & 2 deletions DashAI/back/explainability/explainers/kernel_shap.py
Original file line number Diff line number Diff line change
Expand Up @@ -333,7 +333,10 @@ def fit(
"""
sample_background_data = bool(sample_background_data)

from DashAI.back.explainability.model_input import prepare_model_input
from DashAI.back.explainability.model_input import (
as_shap_predictor,
prepare_model_input,
)

x, y = background_dataset

Expand Down Expand Up @@ -365,7 +368,7 @@ def fit(
import shap

self.explainer = shap.KernelExplainer(
model=self.model.predict,
model=as_shap_predictor(self.model),
data=background_data,
feature_names=feature_names,
link=self.link,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -179,7 +179,10 @@ def fit(
"""
import shap

from DashAI.back.explainability.model_input import prepare_model_input
from DashAI.back.explainability.model_input import (
as_shap_predictor,
prepare_model_input,
)

x, y = background_dataset
# SHAP calls the model with perturbed frames, which skip the model
Expand All @@ -195,7 +198,7 @@ def fit(
background_data = shap.sample(background_data, n_samples)

self.explainer = shap.KernelExplainer(
model=self.model.predict,
model=as_shap_predictor(self.model),
data=background_data,
feature_names=feature_names,
)
Expand Down
40 changes: 39 additions & 1 deletion DashAI/back/explainability/model_input.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,12 +14,50 @@
that both live in the same space.
"""

from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, Callable

if TYPE_CHECKING:
from DashAI.back.dataloaders.classes.dashai_dataset import DashAIDataset


def as_shap_predictor(model: Any) -> Callable:
"""Wrap ``model.predict`` so SHAP receives a plain function, not a method.

SHAP suppresses scikit-learn's "X does not have valid feature names"
warning by blanking ``feature_names_in_`` on whatever object the callable
is bound to (``shap.utils._legacy.convert_to_model``). It reaches that
object through ``__self__``, so it only does this when handed a *bound
method*, and it assumes the attribute is writable.

That assumption does not hold for every model DashAI ships: the LightGBM
and XGBoost wrappers inherit ``feature_names_in_`` from their upstream
estimator as a read-only ``property``, so the assignment raises
``AttributeError: property 'feature_names_in_' ... has no setter`` and the
explanation fails before it starts.

Handing over a plain closure instead leaves ``__self__`` absent, so SHAP
skips that step entirely — a function is SHAP's primary documented
interface for ``model``. The only thing lost is the suppression of a
cosmetic scikit-learn warning.

Parameters
----------
model : Any
The trained model being explained.

Returns
-------
Callable
A one-argument function calling ``model.predict`` positionally, the
same way SHAP calls it today.
"""

def predict(x):
return model.predict(x)

return predict


def prepare_model_input(model: Any, dataset: "DashAIDataset") -> "DashAIDataset":
"""Apply the model's own input preprocessing to a dataset.

Expand Down
30 changes: 30 additions & 0 deletions DashAI/back/initial_components.py
Original file line number Diff line number Diff line change
Expand Up @@ -351,14 +351,31 @@

# Units
from DashAI.back.units.apply_converter_unit import ApplyConverterUnit
from DashAI.back.units.build_global_explainer_unit import BuildGlobalExplainerUnit
from DashAI.back.units.build_local_explainer_unit import BuildLocalExplainerUnit
from DashAI.back.units.build_manual_input_unit import BuildManualInputUnit
from DashAI.back.units.build_model_unit import BuildModelUnit
from DashAI.back.units.evaluate_model_unit import EvaluateModelUnit
from DashAI.back.units.fit_converter_unit import FitConverterUnit
from DashAI.back.units.fit_model_unit import FitModelUnit
from DashAI.back.units.generate_global_explanation_unit import (
GenerateGlobalExplanationUnit,
)
from DashAI.back.units.generate_local_explanation_unit import (
GenerateLocalExplanationUnit,
)
from DashAI.back.units.load_dataset_unit import LoadDatasetUnit
from DashAI.back.units.load_run_model_unit import LoadRunModelUnit
from DashAI.back.units.load_trained_model_unit import LoadTrainedModelUnit
from DashAI.back.units.load_training_dataset_unit import LoadTrainingDatasetUnit
from DashAI.back.units.predict_unit import PredictUnit
from DashAI.back.units.prepare_and_split_unit import PrepareAndSplitUnit
from DashAI.back.units.prepare_explanation_data_unit import PrepareExplanationDataUnit
from DashAI.back.units.run_exploration_unit import RunExplorationUnit
from DashAI.back.units.save_dataset_unit import SaveDatasetUnit
from DashAI.back.units.save_exploration_unit import SaveExplorationUnit
from DashAI.back.units.save_model_unit import SaveModelUnit
from DashAI.back.units.save_prediction_unit import SavePredictionUnit
from DashAI.back.units.transform_dataset_unit import TransformDatasetUnit

logging.basicConfig(level=logging.DEBUG)
Expand Down Expand Up @@ -536,6 +553,19 @@ def get_initial_components():
FitConverterUnit,
TransformDatasetUnit,
SaveDatasetUnit,
RunExplorationUnit,
SaveExplorationUnit,
LoadTrainedModelUnit,
LoadTrainingDatasetUnit,
BuildManualInputUnit,
PredictUnit,
SavePredictionUnit,
LoadRunModelUnit,
BuildGlobalExplainerUnit,
BuildLocalExplainerUnit,
PrepareExplanationDataUnit,
GenerateGlobalExplanationUnit,
GenerateLocalExplanationUnit,
# Explainers
ContrastiveShap,
DiceCounterfactual,
Expand Down
Loading
Loading