Skip to content
Merged
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: 2 additions & 0 deletions .agents/project_context.md
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@ Develop, benchmark, and visualize Explainable AI (xAI) techniques for characteri
- [x] Modular optional dependency extras (`vis`, `drift`, `clustering`, `xai`, `deeplearning`, `dashboard`, `dev`, `all`).
- [x] Single Responsibility Principle (SRP) cleanup, decomposition of monolithic God Classes (`ClusterBasedDriftDetector`), and headless plotting execution without `plt.show()`.
- [x] Standardized GitHub agent instructions, Git & PR workflow skill (`.agents/skills/git-pr-workflow/`), issue templates, and PR template aligned with institutional engineering standards.
- [x] Decouple optional dependency imports (`umap-learn`, `shap`, `lime`, `pyclustering`, `hdbscan`, `tensorflow`, `matplotlib`, `seaborn`) and establish minimal base installation guardrails (`OptionalDependencyError`).
- [ ] PyPI Packaging, Metadata Standardization, and Automated OIDC Trusted Publishing Workflow (#10).
- [ ] Expand automated test coverage for core xAI algorithms in `src/stride/`.
- [ ] Implement additional statistical drift detectors and recurring concept benchmarks.

Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ dependencies = [
"pandas",
"scipy",
"scikit-learn",
"river",
]

[project.optional-dependencies]
Expand Down
11 changes: 10 additions & 1 deletion src/stride/common/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,8 @@
from sklearn.discriminant_analysis import LinearDiscriminantAnalysis as LDA
from sklearn.manifold import MDS, TSNE, LocallyLinearEmbedding
from sklearn.preprocessing import MinMaxScaler, StandardScaler
from umap import UMAP

from stride.exceptions import OptionalDependencyError


class ScalingType(Enum):
Expand Down Expand Up @@ -103,6 +104,14 @@ def _create_reducer(self) -> Any:
elif self.reducer_type == ReducerType.TSNE:
return TSNE(n_components=self.n_components, init="pca", learning_rate="auto", random_state=42)
elif self.reducer_type == ReducerType.UMAP:
try:
from umap import UMAP
except ImportError as err:
raise OptionalDependencyError(
package_name="umap-learn",
feature_name="UMAP dimensionality reduction",
extra_name="clustering",
) from err
return UMAP(n_components=self.n_components, random_state=42, transform_seed=42)
elif self.reducer_type == ReducerType.LLE:
return LocallyLinearEmbedding(n_components=self.n_components, n_neighbors=max(5, self.n_components + 1))
Expand Down
43 changes: 25 additions & 18 deletions src/stride/datasets/protree_data/__init__.py
Original file line number Diff line number Diff line change
@@ -1,25 +1,32 @@
from __future__ import annotations

import click
try:
import click
except ImportError:
click = None

from stride.datasets.protree_data.static import download_all, DEFAULT_DATA_DIR


@click.command()
@click.option("--directory", "-d", default=DEFAULT_DATA_DIR, help="Directory to store datasets")
@click.option("--silent", "-s", is_flag=True, help="Suppress displaying progress.")
@click.option(
"--dataset-names",
"-n",
default="all",
help="Comma-separated list of dataset names to download. "
"Allowable values are 'breast_cancer', 'caltech', 'compass', "
"'diabetes', 'mnist' and 'rhc'. Use 'all' to download all "
"datasets.",
)
def main(directory, silent, dataset_names):
download_all(directory=directory, dataset_names=[s.strip() for s in dataset_names.split(",")], verbose=not silent)
def _cli_entry():
if click is None:
raise ImportError("Dataset download CLI requires 'click'. Install with: pip install click")

@click.command()
@click.option("--directory", "-d", default=DEFAULT_DATA_DIR, help="Directory to store datasets")
@click.option("--silent", "-s", is_flag=True, help="Suppress displaying progress.")
@click.option(
"--dataset-names",
"-n",
default="all",
help="Comma-separated list of dataset names to download. "
"Allowable values are 'breast_cancer', 'caltech', 'compass', "
"'diabetes', 'mnist' and 'rhc'. Use 'all' to download all "
"datasets.",
)
def main(directory, silent, dataset_names):
download_all(directory=directory, dataset_names=[s.strip() for s in dataset_names.split(",")], verbose=not silent)

return main()


if __name__ == "__main__":
main()
_cli_entry()
26 changes: 19 additions & 7 deletions src/stride/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,25 @@ class StrideError(Exception):
class OptionalDependencyError(StrideError):
"""Raised when an optional dependency (e.g. tensorflow, shap) is missing."""

def __init__(self, package_name: str, feature_name: str):
super().__init__(
f"Feature '{feature_name}' requires optional dependency '{package_name}'. "
f"Install it using: pip install stride-xai[{package_name}] or pip install {package_name}"
)
self.package_name = package_name
self.feature_name = feature_name
def __init__(
self,
package_name: str,
feature_name: str | None = None,
extra_name: str | None = None,
):
if feature_name is None:
super().__init__(package_name)
self.package_name = package_name
self.feature_name = ""
self.extra_name = ""
else:
self.package_name = package_name
self.feature_name = feature_name
self.extra_name = extra_name or package_name
super().__init__(
f"Feature '{feature_name}' requires optional dependency '{package_name}'. "
f"Install it using: pip install stride-xai[{self.extra_name}] or pip install {package_name}"
)


class DriftDetectionError(StrideError):
Expand Down
8 changes: 7 additions & 1 deletion src/stride/plotting/_renderers.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,17 @@
from __future__ import annotations

"""Low-level matplotlib rendering primitives for stream visualisation.

These are private helpers used internally by :mod:`.stream`. They are not
part of the public ``stride.plotting`` API.
"""

import numpy as np
import matplotlib.pyplot as plt

try:
import matplotlib.pyplot as plt
except ImportError:
plt = None


def _plot_violin(ax: plt.Axes, values: list, positions: list, colors: list, alphas: list) -> None:
Expand Down
28 changes: 26 additions & 2 deletions src/stride/plotting/stream.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,35 @@
- :func:`visualize_data_stream`
"""

from typing import Any
import numpy as np
import matplotlib.pyplot as plt
import pandas as pd
from matplotlib.figure import Figure
from sklearn.decomposition import PCA

from stride.exceptions import OptionalDependencyError

try:
import matplotlib.pyplot as plt
from matplotlib.figure import Figure

_HAS_MATPLOTLIB = True
except ImportError:
plt = None
Figure = Any # type: ignore
_HAS_MATPLOTLIB = False

from ._renderers import _plot_distribution_comparison


def _ensure_matplotlib() -> None:
if not _HAS_MATPLOTLIB:
raise OptionalDependencyError(
package_name="matplotlib",
feature_name="Stream plotting",
extra_name="vis",
)


def plot_feature_target_relationship(
X,
n_features,
Expand Down Expand Up @@ -60,6 +80,7 @@ def plot_feature_target_relationship(
-------
matplotlib.figure.Figure
"""
_ensure_matplotlib()
unique_classes = sorted(np.unique(np.concatenate([y_before, y_after])))
n_classes = len(unique_classes)

Expand Down Expand Up @@ -122,6 +143,7 @@ def plot_class_distribution(class_dist_before, class_dist_after, class_colors, t
-------
matplotlib.figure.Figure
"""
_ensure_matplotlib()
fig, (ax_before, ax_after) = plt.subplots(1, 2, figsize=(12, 6))
if title:
fig.suptitle(title, fontsize=16, fontweight="bold", y=1.0)
Expand Down Expand Up @@ -175,6 +197,7 @@ def plot_feature_space(
-------
matplotlib.figure.Figure
"""
_ensure_matplotlib()
fig, (ax_before, ax_after) = plt.subplots(1, 2, figsize=(14, 7))
fs_title_suffix = ""

Expand Down Expand Up @@ -307,6 +330,7 @@ def visualize_data_stream(
list[matplotlib.figure.Figure]
List of three figures in the order described above.
"""
_ensure_matplotlib()
if isinstance(X, pd.DataFrame):
X = X.values
if isinstance(y, pd.Series):
Expand Down
16 changes: 9 additions & 7 deletions src/stride/xai/boundary/analysis.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import numpy as np
import random
from sklearn.preprocessing import MinMaxScaler
from stride.exceptions import OptionalDependencyError, StrideError
from stride.xai.boundary.disagreement import compute_disagreement_analysis


Expand Down Expand Up @@ -64,14 +65,15 @@ def analyze(self, model_class=None, model_params=None, grid_size=300, ssnp_epoch
else:
try:
from stride.xai.boundary.ssnp import SSNP
except ImportError as err:
raise ImportError(
"High-dimensional decision boundary projection requires the deeplearning extra: "
"install with 'pip install stride-xai[deeplearning]' (requires tensorflow)."

ssnp = SSNP(epochs=ssnp_epochs, patience=ssnp_patience, verbose=0)
ssnp.fit(X_before_scaled, self.y_before)
except Exception as err:
raise OptionalDependencyError(
package_name="tensorflow",
feature_name="High-dimensional decision boundary projection (SSNP)",
extra_name="deeplearning",
) from err
# SSNP is used to find a 2D projection that preserves class structure.
ssnp = SSNP(epochs=ssnp_epochs, patience=ssnp_patience, verbose=0)
ssnp.fit(X_before_scaled, self.y_before)

# Project points to 2D (if 2D already, this just returns the scaled data)
X_before_2d = ssnp.transform(X_before_scaled)
Expand Down
28 changes: 22 additions & 6 deletions src/stride/xai/boundary/ssnp.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,21 @@
import os
import numpy as np
from sklearn.preprocessing import LabelBinarizer
import tensorflow as tf
from tensorflow.keras import regularizers
from tensorflow.keras.callbacks import EarlyStopping
from tensorflow.keras.initializers import Constant
from tensorflow.keras.layers import Dense, Input
from tensorflow.keras.models import Model

from stride.exceptions import OptionalDependencyError

try:
import tensorflow as tf
from tensorflow.keras import regularizers
from tensorflow.keras.callbacks import EarlyStopping
from tensorflow.keras.initializers import Constant
from tensorflow.keras.layers import Dense, Input
from tensorflow.keras.models import Model

_HAS_TF = True
except ImportError:
_HAS_TF = False
tf = None

# Ensure deterministic operations where possible
os.environ["TF_DETERMINISTIC_OPS"] = "1"
Expand Down Expand Up @@ -54,6 +63,13 @@ def __init__(
self.inv = None
self.clustering = None

if not _HAS_TF:
raise OptionalDependencyError(
package_name="tensorflow",
feature_name="SSNP boundary projection",
extra_name="deeplearning",
)

tf.random.set_seed(42)

tf.keras.backend.clear_session()
Expand Down
16 changes: 10 additions & 6 deletions src/stride/xai/clustering/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
from .clustering import ClusterBasedDriftDetector # noqa: F401
from .visualization import (
plot_drift_clustered,
plot_clusters_by_class,
plot_centers_shift,
plot_clustering_heatmap,
) # noqa: F401

try:
from .visualization import (
plot_drift_clustered,
plot_clusters_by_class,
plot_centers_shift,
plot_clustering_heatmap,
) # noqa: F401
except ImportError:
pass
11 changes: 10 additions & 1 deletion src/stride/xai/clustering/xmeans.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from typing import Sequence

import numpy as np
from pyclustering.cluster.xmeans import kmeans_plusplus_initializer, xmeans # type: ignore
from stride.exceptions import OptionalDependencyError


def reshape_clusters(clusters: Sequence[Sequence[int]]) -> np.ndarray:
Expand Down Expand Up @@ -63,6 +63,15 @@ def run_xmeans(
here seeds both ``numpy.random`` and the built-in ``random`` module, which is
sufficient for pyclustering's internal sampling.
"""
try:
from pyclustering.cluster.xmeans import kmeans_plusplus_initializer, xmeans
except ImportError as err:
raise OptionalDependencyError(
package_name="pyclustering",
feature_name="X-Means clustering",
extra_name="clustering",
) from err

if random_state is not None:
random.seed(random_state)
np.random.seed(random_state)
Expand Down
12 changes: 8 additions & 4 deletions src/stride/xai/importance/__init__.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
from .base import FeatureImportanceMethod # noqa: F401
from .methods import calculate_feature_importance # noqa: F401
from .visualization import ( # noqa: F401
visualize_drift_importance,
visualize_predictive_importance_shift,
)

try:
from .visualization import ( # noqa: F401
visualize_drift_importance,
visualize_predictive_importance_shift,
)
except ImportError:
pass
from .analysis import FeatureImportanceDriftAnalyzer # noqa: F401
21 changes: 19 additions & 2 deletions src/stride/xai/importance/methods.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
import numpy as np
from sklearn.inspection import permutation_importance
import shap
from lime.lime_tabular import LimeTabularExplainer
from stride.exceptions import OptionalDependencyError
from .base import FeatureImportanceMethod


Expand Down Expand Up @@ -68,6 +67,15 @@ def _calculate_pfi(model, X, y, n_repeats=30, random_state=42):

def _calculate_shap(model, X, feature_names):
"""Calculate SHAP values."""
try:
import shap
except ImportError as err:
raise OptionalDependencyError(
package_name="shap",
feature_name="SHAP feature importance",
extra_name="xai",
) from err

# Use a subset for efficiency if dataset is large
background_size = min(100, len(X))
background = shap.sample(X, background_size)
Expand Down Expand Up @@ -123,6 +131,15 @@ def _calculate_shap(model, X, feature_names):

def _calculate_lime(model, X, y, feature_names, random_state=42):
"""Calculate LIME feature importance."""
try:
from lime.lime_tabular import LimeTabularExplainer
except ImportError as err:
raise OptionalDependencyError(
package_name="lime",
feature_name="LIME feature importance",
extra_name="xai",
) from err

np.random.seed(random_state)

# Create LIME explainer
Expand Down
Loading
Loading