Interprétabilité - Explicateur SHAP tabulaire

Utilisez Kernel SHAP (SHapley Additive exPlanations) pour expliquer un modèle de classification tabulaire. Noyau SHAP est une méthode indépendante du modèle qui estime la contribution de chaque fonctionnalité à la prédiction d’un modèle. Vous entraînez un modèle de régression logistique sur le jeu de données Adult Census Income, puis utilisez le transformateur SynapseML TabularSHAP pour calculer des explications au niveau des caractéristiques.

Prerequisites

  • Créez un notebook dans votre espace de travail et rattachez-le à un lakehouse. Pour plus d’informations, consultez Créer un bloc-notes.

SynapseML, PySpark, pandas et plotly sont préinstallés dans les environnements de notebooks Fabric. Aucune installation supplémentaire du package n’est requise.

Importer des packages et définir des fonctions utilitaires définies par l’utilisateur

Dans votre bloc-notes Fabric, collez le code suivant dans une cellule et exécutez-le. Cette étape importe les bibliothèques requises et définit deux fonctions définies par l’utilisateur pour extraire les éléments vectoriels ultérieurement.

import pyspark
from synapse.ml.explainers import TabularSHAP
from pyspark.ml import Pipeline
from pyspark.ml.classification import LogisticRegression
from pyspark.ml.feature import StringIndexer, OneHotEncoder, VectorAssembler
from pyspark.sql.types import FloatType, ArrayType
from pyspark.sql.functions import col, lit, rand, broadcast, udf
import pandas as pd

vec_access = udf(lambda v, i: float(v[i]), FloatType())
vec2array = udf(lambda vec: vec.toArray().tolist(), ArrayType(FloatType()))

Vérifiez : exécutez le code suivant dans une nouvelle cellule. Vous devriez voir le résultat TabularSHAP imported successfully.

print("TabularSHAP imported successfully")
print(f"PySpark version: {pyspark.__version__}")

Charger des données et entraîner un modèle de classification

Chargez le jeu de données Adult Census Income depuis le stockage Blob Azure, indexez le libellé cible et entraînez un pipeline de régression logistique.

df = spark.read.parquet(
    "wasbs://publicwasb@mmlspark.blob.core.windows.net/AdultCensusIncome.parquet"
)

labelIndexer = StringIndexer(
    inputCol="income", outputCol="label", stringOrderType="alphabetAsc"
).fit(df)
print("Label index assignment: " + str(set(zip(labelIndexer.labels, [0, 1]))))

training = labelIndexer.transform(df).cache()

categorical_features = [
    "workclass",
    "education",
    "marital-status",
    "occupation",
    "relationship",
    "race",
    "sex",
    "native-country",
]
categorical_features_idx = [feat + "_idx" for feat in categorical_features]
categorical_features_enc = [feat + "_enc" for feat in categorical_features]
numeric_features = [
    "age",
    "education-num",
    "capital-gain",
    "capital-loss",
    "hours-per-week",
]

strIndexer = StringIndexer(
    inputCols=categorical_features, outputCols=categorical_features_idx
)
onehotEnc = OneHotEncoder(
    inputCols=categorical_features_idx, outputCols=categorical_features_enc
)
vectAssem = VectorAssembler(
    inputCols=categorical_features_enc + numeric_features, outputCol="features"
)
lr = LogisticRegression(featuresCol="features", labelCol="label", weightCol="fnlwgt")
pipeline = Pipeline(stages=[strIndexer, onehotEnc, vectAssem, lr])
model = pipeline.fit(training)

Vérifiez : exécutez la cellule suivante. Vous devez voir les nombres de lignes pour les données d’apprentissage et la confirmation des étapes du pipeline.

print(f"Training rows: {training.count()}")
print(f"Pipeline stages: {[type(s).__name__ for s in model.stages]}")
assert training.count() > 30000, "Dataset should contain over 30,000 rows"
print("Model trained successfully")

# Expected output:
#Training rows: 32561
#Pipeline stages: ['StringIndexerModel', 'OneHotEncoderModel', #'VectorAssembler', 'LogisticRegressionModel']
#Model trained successfully

Sélectionner des observations à expliquer

Sélectionnez de façon aléatoire cinq observations dans les données d’apprentissage notées. Ces observations sont les instances pour lesquelles vous générez des explications SHAP.

explain_instances = (
    model.transform(training).orderBy(rand()).limit(5).repartition(200).cache()
)
display(explain_instances)

Vérifiez : confirmez la taille de l’échantillon.

count = explain_instances.count()
print(f"Explain instances: {count}")
assert count == 5, f"Expected 5 rows, got {count}"
print("Sample selected successfully")

Configurer et exécuter TabularSHAP

Créez un TabularSHAP explicatif et appliquez-le aux observations sélectionnées. Les paramètres clés sont les suivants :

Paramètre Description
inputCols Colonnes de caractéristiques que le modèle utilise pour la prédiction.
outputCol Nom de la colonne qui contient des valeurs de sortie SHAP.
numSamples Nombre d’échantillons de perturbation pour l’estimation SHAP du noyau. Les valeurs plus élevées sont plus précises mais plus lentes.
model Modèle de pipeline entraîné à expliquer.
targetCol Colonne de sortie du modèle à expliquer. Dans cet exemple, la colonne est probability.
targetClasses Indices de classe à expliquer. [1] explique seulement la probabilité de la classe 1. Utilisez [0, 1] pour expliquer les deux classes.
backgroundData Exemple de données d’apprentissage utilisées comme distribution de référence pour l’intégration des fonctionnalités.
shap = TabularSHAP(
    inputCols=categorical_features + numeric_features,
    outputCol="shapValues",
    numSamples=5000,
    model=model,
    targetCol="probability",
    targetClasses=[1],
    backgroundData=broadcast(training.orderBy(rand()).limit(100).cache()),
)

shap_df = shap.transform(explain_instances)

Note

Cette étape peut prendre plusieurs minutes en fonction numSamples de la taille du cluster. Avec numSamples=5000 et cinq observations, prévoyez 3 à 10 minutes sur un cluster Spark Fabric par défaut.

Vérifiez : vérifiez que la colonne de sortie SHAP existe.

assert "shapValues" in shap_df.columns, "shapValues column missing"
print(f"SHAP output columns: {shap_df.columns}")
print("TabularSHAP transform completed")

Extraire les valeurs SHAP

Extrayez la probabilité de la classe 1 et les valeurs SHAP du DataFrame de résultats. Pour chaque observation, le vecteur de valeurs SHAP commence par la valeur de base (sortie moyenne du jeu de données d’arrière-plan), suivie d’une valeur par fonctionnalité.

shaps = (
    shap_df.withColumn("probability", vec_access(col("probability"), lit(1)))
    .withColumn("shapValues", vec2array(col("shapValues").getItem(0)))
    .select(
        ["shapValues", "probability", "label"] + categorical_features + numeric_features
    )
)

shaps_local = shaps.toPandas()
shaps_local.sort_values("probability", ascending=False, inplace=True, ignore_index=True)
pd.set_option("display.max_colwidth", None)
display(shaps_local)

Vérifiez : confirmez la structure du DataFrame pandas.

expected_cols = len(categorical_features) + len(numeric_features) + 3
print(f"DataFrame shape: {shaps_local.shape}")
print(f"Expected columns: {expected_cols}, Actual: {shaps_local.shape[1]}")
assert shaps_local.shape == (5, expected_cols), f"Unexpected shape: {shaps_local.shape}"
print("SHAP values extracted successfully")

Visualiser les valeurs SHAP

Créez un graphique à barres pour chaque observation qui montre comment chaque fonctionnalité contribue à la probabilité prédite.

from plotly.subplots import make_subplots
import plotly.graph_objects as go

features = categorical_features + numeric_features
features_with_base = ["Base"] + features

rows = shaps_local.shape[0]

fig = make_subplots(
    rows=rows,
    cols=1,
    subplot_titles="Probability: "
    + shaps_local["probability"].apply("{:.2%}".format)
    + "; Label: "
    + shaps_local["label"].astype(str),
)

for index, row in shaps_local.iterrows():
    feature_values = [0] + [row[feature] for feature in features]
    shap_values = row["shapValues"]
    list_of_tuples = list(zip(features_with_base, feature_values, shap_values))
    shap_pdf = pd.DataFrame(list_of_tuples, columns=["name", "value", "shap"])
    fig.add_trace(
        go.Bar(
            x=shap_pdf["name"],
            y=shap_pdf["shap"],
            hovertext="value: " + shap_pdf["value"].astype(str),
        ),
        row=index + 1,
        col=1,
    )

fig.update_yaxes(range=[-1, 1], fixedrange=True, zerolinecolor="black")
fig.update_xaxes(type="category", tickangle=45, fixedrange=True)
fig.update_layout(height=400 * rows, title_text="SHAP explanations")
fig.show()

Vérifiez : confirmez que l’objet de graphique a été créé.

print(f"Figure traces: {len(fig.data)}")
print(f"Figure height: {fig.layout.height}px")
assert len(fig.data) == 5, f"Expected 5 traces, got {len(fig.data)}"
print("Visualization created successfully")

Interpréter les résultats

Chaque sous-diagramme représente une observation. Les barres indiquent :

  • Base : sortie moyenne du modèle dans le jeu de données en arrière-plan (probabilité de référence).
  • Valeurs SHAP positives : fonctionnalités qui poussent la prédiction vers la classe 1 (revenu supérieur à 50 000).
  • Valeurs SHAP négatives : caractéristiques qui poussent la prédiction vers la classe 0 (revenu inférieur ou égal à 50 000).

La somme de la valeur de base et de toutes les valeurs SHAP des caractéristiques est égale à la probabilité prédite du modèle pour cette observation.

Résolution des problèmes

Problème Cause Résolution
OutOfMemoryError lors de l’exécution de TabularSHAP numSamples est trop volumineux pour la mémoire disponible. Réduisez , par exemple à 1 000, ou augmentez numSamplesla mémoire de l’exécuteur Spark.
La transformation SHAP est lente Une valeur élevée de numSamples avec de nombreuses fonctions augmente le temps de calcul. Réduisez numSamples à 1 000-2 000 pour obtenir des résultats exploratoires plus rapides. Augmentation pour l’analyse finale.
FileNotFoundException pour parquet L’accès réseau à mmlspark.blob.core.windows.net est bloqué. Vérifiez que votre espace de travail Fabric dispose d’un accès Internet sortant. Vous pouvez également charger le jeu de données dans votre lakehouse.
shapValues colonne contient des valeurs Null Certaines observations peuvent ne pas être correctement traitées si les valeurs des variables se situent en dehors de la distribution des données d’entraînement. Recherchez les valeurs null ou inattendues dans les fonctionnalités d’entrée. Filtrez les valeurs Null des résultats.
display() n’affiche aucune sortie Le code s’exécute en dehors d’un environnement de notebook Fabric. Utilisez shaps_local.head() ou print(shaps_local) dans des environnements de Python standard.

Nettoyage

Si vous avez chargé le jeu de données dans un lakehouse pour ce didacticiel, supprimez-le pour libérer le stockage :

# Remove cached DataFrames from memory
training.unpersist()
explain_instances.unpersist()
print("Cached DataFrames released")