Notitie
Voor toegang tot deze pagina is autorisatie vereist. U kunt proberen u aan te melden of de directory te wijzigen.
Voor toegang tot deze pagina is autorisatie vereist. U kunt proberen de mappen te wijzigen.
Gebruik Kernel SHAP (SHapley Additive exPlanations) om een tabellair classificatiemodel uit te leggen. Kernel SHAP is een modelagnostische methode waarmee de bijdrage van elke functie aan de voorspelling van een model wordt geschat. U traint een logistiek regressiemodel op de gegevensset Adult Census Income en gebruikt vervolgens de SynapseML-transformator TabularSHAP om uitleg op functieniveau te berekenen.
Prerequisites
Haal een Microsoft Fabric-abonnement op. Of meld u aan voor een gratis proefversie van Microsoft Fabric.
Meld u aan bij Microsoft Fabric.
Schakel over naar Fabric met behulp van de ervaringsschakelaar aan de linkerkant van de startpagina.
- Maak een nieuw notitieblok in uw werkruimte en koppel dit aan een lakehouse. Zie Een notitieblok maken voor meer informatie.
SynapseML, PySpark, pandas en plotly zijn vooraf geïnstalleerd in Fabric notebookomgevingen. Er is geen extra pakketinstallatie vereist.
Pakketten importeren en hulp-UDF's definiëren
Plak in uw Fabric notebook de volgende code in een cel en voer deze uit. Met deze stap importeert u de vereiste bibliotheken en definieert u twee door de gebruiker gedefinieerde functies (UDF's) om later vectorelementen te extraheren.
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()))
Controleer: Voer de volgende code uit in een nieuwe cel. U ziet nu de uitvoer TabularSHAP imported successfully.
print("TabularSHAP imported successfully")
print(f"PySpark version: {pyspark.__version__}")
Gegevens laden en een classificatiemodel trainen
Laad de gegevensset Adult Census Income van Azure Blob Storage, indexeer het doellabel en train een logistieke regressiepijplijn.
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)
Controleer: Voer de volgende cel uit. U ziet het aantal rijen voor trainingsgegevens en de bevestiging van pijplijnfasen.
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
Observaties selecteren om uit te leggen
Selecteer willekeurig vijf waarnemingen uit de gescoorde trainingsgegevens. Deze waarnemingen zijn de exemplaren waarvoor u SHAP-uitleg genereert.
explain_instances = (
model.transform(training).orderBy(rand()).limit(5).repartition(200).cache()
)
display(explain_instances)
Controleer: Bevestig de grootte van het voorbeeld.
count = explain_instances.count()
print(f"Explain instances: {count}")
assert count == 5, f"Expected 5 rows, got {count}"
print("Sample selected successfully")
TabularSHAP configureren en uitvoeren
Maak een TabularSHAP uitleg en pas deze toe op de geselecteerde waarnemingen. De belangrijkste parameters zijn:
| Parameter | Description |
|---|---|
inputCols |
Functiekolommen die door het model worden gebruikt voor voorspelling. |
outputCol |
Naam van de kolom die SHAP-uitvoerwaarden bevat. |
numSamples |
Aantal verstoringsvoorbeelden voor kernel SHAP-schatting. Hogere waarden zijn nauwkeuriger, maar langzamer. |
model |
Het getrainde pijplijnmodel om uit te leggen. |
targetCol |
De uit te leggen uitvoerkolom van het model. In dit voorbeeld is de kolom probability. |
targetClasses |
Klassenindexen die moeten worden uitgelegd.
[1] verklaart alleen kans van klasse 1. Gebruik [0, 1] dit om beide klassen uit te leggen. |
backgroundData |
Een voorbeeld van trainingsgegevens die worden gebruikt als referentiedistributie voor het integreren van functies. |
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
Deze stap kan enkele minuten duren, afhankelijk numSamples van en clustergrootte. Met numSamples=5000 en vijf waarnemingen verwacht u 3-10 minuten op een standaard Fabric Spark-cluster.
Controleer: Controleer of de SHAP-uitvoerkolom bestaat.
assert "shapValues" in shap_df.columns, "shapValues column missing"
print(f"SHAP output columns: {shap_df.columns}")
print("TabularSHAP transform completed")
SHAP-waarden extraheren
Pak de waarschijnlijkheids- en SHAP-waarden van klasse 1 uit het dataframe van het resultaat. Voor elke observatie begint de SHAP-waardenvector met de basiswaarde (gemiddelde uitvoer van de achtergrondgegevensset), gevolgd door één waarde per functie.
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)
Controleer: Bevestig de pandas DataFrame-structuur.
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")
SHAP-waarden visualiseren
Maak een staafdiagram voor elke observatie die laat zien hoe elke functie bijdraagt aan de voorspelde waarschijnlijkheid.
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()
Verifiëren: Bevestig dat het plotobject is aangemaakt.
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")
De resultaten interpreteren
Elk subplot vertegenwoordigt één observatie. De balken geven het volgende weer:
- Basis: de gemiddelde modeluitvoer voor de achtergrondgegevensset (basislijnkans).
- Positieve SHAP-waarden: functies die de voorspelling naar klasse 1 pushen (inkomsten groter dan 50.000).
- Negatieve SHAP-waarden: functies die de voorspelling naar klasse 0 pushen (inkomen kleiner dan of gelijk aan 50K).
De som van de basiswaarde en alle functie-SHAP-waarden is gelijk aan de voorspelde waarschijnlijkheid van het model voor die observatie.
Troubleshooting
| Probleem | Oorzaak | Resolutie |
|---|---|---|
OutOfMemoryError tijdens TabularSHAP |
numSamples is te groot voor beschikbaar geheugen. |
Verminder numSamplesbijvoorbeeld tot 1.000 of verhoog het Spark-uitvoerprogrammageheugen. |
| SHAP-transformatie is traag | Hoog numSamples met veel functies verhoogt de rekentijd. |
Verminder numSamples tot 1.000-2.000 voor snellere verkennende resultaten. Verhoging voor uiteindelijke analyse. |
FileNotFoundException voor parket |
Netwerktoegang tot mmlspark.blob.core.windows.net is geblokkeerd. |
Controleer of uw Fabric werkruimte uitgaande internettoegang heeft. U kunt de gegevensset ook uploaden naar uw lakehouse. |
shapValues kolom bevat null-waarden |
Sommige waarnemingen kunnen mislukken als functiewaarden buiten de trainingsdistributie vallen. | Controleer op null- of onverwachte waarden in invoerfuncties. Null-waarden filteren op basis van resultaten. |
display() geeft geen uitvoer weer |
De code wordt uitgevoerd buiten een Fabric notebookomgeving. | Gebruik shaps_local.head() of print(shaps_local) in standaardomgevingen Python. |
Schoonmaken
Als u de dataset voor deze zelfstudie naar een lakehouse hebt geüpload, verwijder deze dan om opslagruimte vrij te maken:
# Remove cached DataFrames from memory
training.unpersist()
explain_instances.unpersist()
print("Cached DataFrames released")