Możliwość interpretacji — tabelaryczny objaśnienie SHAP

Użyj metody Kernel SHAP (SHapley Additive exPlanations), aby objaśnić model klasyfikacji danych tabelarycznych. Kernel SHAP to niezależna od modelu metoda, która szacuje wkład każdej cechy w predykcję modelu. Wytrenujesz model regresji logistycznej w zestawie danych Adult Census Income, a następnie użyjesz transformatora SynapseML TabularSHAP , aby obliczyć wyjaśnienia na poziomie funkcji.

Prerequisites

  • Uzyskaj subskrypcję usługi Microsoft Fabric. Możesz też utworzyć konto bezpłatnej wersji próbnej usługi Microsoft Fabric.

  • Zaloguj się do Microsoft Fabric.

  • Przełącz na Fabric, używając przełącznika doświadczenia w dolnym lewym rogu twojej strony głównej.

    Zrzut ekranu przedstawiający wybór Fabric w menu przełącznika środowiska.

  • Utwórz nowy notesnik w obszarze roboczym i dołącz go do magazynu danych. Aby uzyskać więcej informacji, zobacz Tworzenie notesu.

SynapseML, PySpark, pandas i plotly są preinstalowane w środowiskach notesników Fabric. Nie jest wymagana dodatkowa instalacja pakietu.

Importowanie pakietów i definiowanie pomocniczych funkcji UDF

W notesie Fabric wklej następujący kod do komórki i uruchom go. Ten krok importuje wymagane biblioteki i definiuje dwie funkcje zdefiniowane przez użytkownika (UDF) do wyodrębniania elementów wektorów później.

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()))

Sprawdź: uruchom następujący kod w nowej komórce. Powinny zostać wyświetlone dane wyjściowe TabularSHAP imported successfully.

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

Ładowanie danych i trenowanie modelu klasyfikacji

Załaduj zestaw danych Adult Census Income z Azure Blob Storage, zaindeksuj etykietę docelową i wytrenuj potok regresji logistycznej.

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)

Sprawdź: Uruchom następującą komórkę. Powinny zostać wyświetlone liczby wierszy dla danych szkoleniowych i potwierdzenie etapów potoku.

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

Wybierz obserwacje, aby wyjaśnić

Losowo wybierz pięć obserwacji z ocenianych danych treningowych. Te obserwacje to wystąpienia, dla których generujesz wyjaśnienia SHAP.

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

Sprawdź: Potwierdź rozmiar próbki.

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

Konfigurowanie i uruchamianie programu TabularSHAP

Utwórz objaśnienie TabularSHAP i zastosuj je do wybranych obserwacji. Kluczowe parametry to:

Parameter Opis
inputCols Kolumny funkcji używane przez model do przewidywania.
outputCol Nazwa kolumny zawierającej wartości wyjściowe SHAP.
numSamples Liczba próbek perturbacji do estymacji metodą Kernel SHAP. Wyższe wartości są dokładniejsze, ale wolniejsze.
model Wytrenowany model potoku do wyjaśnienia.
targetCol Kolumna danych wyjściowych modelu do wyjaśnienia. W tym przykładzie kolumna to probability.
targetClasses Indeksy klas do objaśnienia. [1] wyjaśnia tylko prawdopodobieństwo klasy 1. Użyj [0, 1], aby wyjaśnić obie klasy.
backgroundData Przykład danych szkoleniowych używanych jako dystrybucja referencyjna do integrowania funkcji.
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

Ten krok może potrwać kilka minut w zależności od numSamples i rozmiaru klastra. W przypadku numSamples=5000 i pięciu obserwacji spodziewaj się 3–10 minut w domyślnym klastrze Fabric Spark.

Sprawdź: Sprawdź, czy kolumna danych wyjściowych SHAP istnieje.

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

Wyodrębnij wartości SHAP

Wyodrębnij wartości prawdopodobieństwa klasy 1 i SHAP z wynikowej ramki danych. Dla każdej obserwacji wektor wartości SHAP rozpoczyna się od wartości podstawowej (średniej danych wyjściowych zestawu danych w tle), a następnie jednej wartości na funkcję.

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)

Sprawdź: Potwierdź strukturę obiektu DataFrame biblioteki 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")

Wizualizowanie wartości SHAP

Utwórz wykres słupkowy dla każdej obserwacji, który pokazuje, jak każda funkcja przyczynia się do przewidywanego prawdopodobieństwa.

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()

Sprawdź: Upewnij się, że obiekt wykresu został utworzony.

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")

Interpretacja wyników

Każdy podlot reprezentuje jedną obserwację. Słupki pokazują:

  • Podstawa: średnie dane wyjściowe modelu w zestawie danych w tle (prawdopodobieństwo punktu odniesienia).
  • Dodatnie wartości SHAP: cechy, które przesuwają predykcję w kierunku klasy 1 (dochód większy niż 50 tys.).
  • Ujemne wartości SHAP: funkcje, które wypychają przewidywanie do klasy 0 (dochód mniejszy lub równy 50K).

Suma wartości bazowej i wszystkich wartości SHAP cech jest równa przewidywanemu przez model prawdopodobieństwu dla tej obserwacji.

Troubleshooting

Problematyka Przyczyna Resolution
OutOfMemoryError podczas TabularSHAP numSamples jest za duży dla dostępnej pamięci. Zmniejsz numSamples, na przykład do 1 000, lub zwiększ pamięć executora Spark.
Transformacja SHAP jest powolna Wysokie numSamples z wieloma funkcjami zwiększa czas obliczeniowy. Zmniejsz numSamples do 1000–2000, aby uzyskać szybsze wyniki eksploracyjne. Zwiększ wartość na potrzeby ostatecznej analizy.
FileNotFoundException dla parquet Dostęp do mmlspark.blob.core.windows.net z sieci jest zablokowany. Sprawdź, czy obszar roboczy Fabric ma wychodzący dostęp do Internetu. Alternatywnie prześlij zestaw danych do swojego lakehouse’u.
shapValues kolumna zawiera wartości null Niektóre obserwacje mogą zakończyć się niepowodzeniem, jeśli wartości funkcji znajdują się poza rozkładem trenowania. Sprawdź, czy w cechach wejściowych występują wartości null lub wartości nieoczekiwane. Filtruj wartości null z wyników.
display() pokazuje brak danych wyjściowych Kod działa poza środowiskiem notesu Fabric. Użyj shaps_local.head() lub print(shaps_local) w standardowych środowiskach Python.

Czyszczenie

Jeśli przesłano zestaw danych do lakehouse w ramach tego samouczka, usuń go, aby zwolnić miejsce w magazynie:

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