Eksplorator jądra SHAP

Wyjaśnialne AI w Pythonie

Fouad Trad

Machine Learning Engineer

Eksplorator jądra SHAP

 

  • Wyznacza wartości SHAP dla dowolnego modelu

    • K najbliższych sąsiadów
    • Sieci neuronowe
    • Modele drzewiaste
  • Wolniejszy niż eksploratory dedykowane

Schemat przedstawiający podział eksploratora SHAP na eksploratory ogólne (dla dowolnych modeli) i dedykowane (zoptymalizowane dla konkretnych typów modeli).

Wyjaśnialne AI w Pythonie

Choroba serca

age sex chest_pain_type blood_pressure ecg_results thalassemia target
52 1 0 125 1 3 0
53 1 0 140 0 3 0
70 1 0 145 1 3 0
61 1 0 148 1 3 0
62 0 0 138 1 2 0

 

mlp_clf: wielowarstwowy perceptron przewidujący ryzyko choroby serca

Wyjaśnialne AI w Pythonie

Koszty ubezpieczenia

age gender bmi children smoker charges
19 0 27.900 0 1 16884.92
18 1 33.770 1 0 1725.55
28 1 33.000 3 0 4449.46
33 1 22.705 0 0 21984.47
32 1 28.880 0 0 3866.85

 

mlp_reg: wielowarstwowy perceptron przewidujący koszty ubezpieczenia

Wyjaśnialne AI w Pythonie

Tworzenie eksploratora jądra

MLPRegressor
import shap


explainer = shap.KernelExplainer( # Funkcja predykcji modelu, # Reprezentatywne podsumowanie zbioru danych )
MLPClassifier
import shap


explainer = shap.KernelExplainer( # Funkcja predykcji modelu, # Reprezentatywne podsumowanie zbioru danych )
Wyjaśnialne AI w Pythonie

Tworzenie eksploratora jądra

MLPRegressor
import shap

explainer = shap.KernelExplainer(
  mlp_reg.predict, 
  # Reprezentatywne podsumowanie zbioru danych
)


MLPClassifier
import shap

explainer = shap.KernelExplainer(
  mlp_clf.predict_proba, 
  # Reprezentatywne podsumowanie zbioru danych
)


Wyjaśnialne AI w Pythonie

Tworzenie eksploratora jądra

MLPRegressor
import shap

explainer = shap.KernelExplainer(
  mlp_reg.predict, 
  shap.kmeans(X, 10)
)


shap_values_reg = explainer.shap_values(X)
MLPClassifier
import shap

explainer = shap.KernelExplainer(
  mlp_clf.predict_proba, 
  shap.kmeans(X, 10)
)


shap_values_cls = explainer.shap_values(X)
Wyjaśnialne AI w Pythonie

Ważność cech

MLPRegressor
mean_reg = np.abs(shap_values_reg).mean(axis=0)

plt.bar(X.columns, mean_reg)

Wykres słupkowy ważności cech w zadaniu regresji – palenie i wiek mają największy wpływ na przewidywane koszty.

MLPClassifier
mean_cls = np.abs(shap_values_cls[:,:,1]).mean(axis=0)

plt.bar(X.columns, mean_cls)

Wykres słupkowy ważności cech w zadaniu klasyfikacji – typ bólu w klatce piersiowej i talasemia mają największy wpływ na przewidywane wyniki.

Wyjaśnialne AI w Pythonie

Porównanie z podejściami dedykowanymi

Regresja liniowa
plt.bar(X.columns, np.abs(lin_reg.coef_))

Wykres słupkowy ważności cech w regresji liniowej – palenie i wiek mają największy wpływ na przewidywane koszty.

Regresja logistyczna
plt.bar(X.columns, np.abs(log_reg.coef_[0]))

Wykres słupkowy ważności cech w zadaniu klasyfikacji – typ bólu w klatce piersiowej ma największy wpływ na przewidywane wyniki.

Wyjaśnialne AI w Pythonie

Ćwiczenia!

Wyjaśnialne AI w Pythonie

Preparing Video For Download...