Пояснювач ядра SHAP

Пояснюваний ШІ в Python

Fouad Trad

Machine Learning Engineer

Пояснювач ядра SHAP

 

  • Обчислює значення SHAP для будь-якої моделі

    • K-nearest neighbors
    • Нейронні мережі
    • Деревоподібні моделі
  • Повільніший за спеціалізовані пояснювачі

Зображення показує, що пояснювачі SHAP поділяються на загальні, які можна застосувати до будь-якої моделі, і спеціалізовані, оптимізовані під окремі типи моделей.

Пояснюваний ШІ в Python

Серцеві хвороби

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: багатошаровий перцептрон для прогнозування ризику серцевих хвороб

Пояснюваний ШІ в Python

Страхові витрати

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: багатошаровий перцептрон для прогнозування страхових витрат

Пояснюваний ШІ в Python

Створення пояснювачів ядра

MLPRegressor
import shap


explainer = shap.KernelExplainer( # Функція передбачення моделі, # Репрезентативна вибірка набору даних )
MLPClassifier
import shap


explainer = shap.KernelExplainer( # Функція передбачення моделі, # Репрезентативна вибірка набору даних )
Пояснюваний ШІ в Python

Створення пояснювачів ядра

MLPRegressor
import shap

explainer = shap.KernelExplainer(
  mlp_reg.predict, 
  # Репрезентативна вибірка набору даних
)


MLPClassifier
import shap

explainer = shap.KernelExplainer(
  mlp_clf.predict_proba, 
  # Репрезентативна вибірка набору даних
)


Пояснюваний ШІ в Python

Створення пояснювачів ядра

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)
Пояснюваний ШІ в Python

Важливість ознак

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

plt.bar(X.columns, mean_reg)

Зображення показує стовпчикову діаграму важливості ознак для задачі регресії: куріння та вік — найвпливовіші фактори у прогнозуванні витрат.

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

plt.bar(X.columns, mean_cls)

Зображення показує стовпчикову діаграму важливості ознак для задачі класифікації: тип болю в грудях і таласемія — найвпливовіші фактори у прогнозуванні витрат.

Пояснюваний ШІ в Python

Порівняння зі спеціалізованими підходами

Лінійна регресія
plt.bar(X.columns, np.abs(lin_reg.coef_))

Зображення показує стовпчикову діаграму важливості ознак у моделі лінійної регресії: куріння та вік — найвпливовіші фактори у прогнозуванні витрат.

Логістична регресія
plt.bar(X.columns, np.abs(log_reg.coef_[0]))

Зображення показує стовпчикову діаграму важливості ознак для задачі класифікації: тип болю в грудях — найвпливовіший фактор у прогнозуванні витрат.

Пояснюваний ШІ в Python

Давайте потренуємось!

Пояснюваний ШІ в Python

Preparing Video For Download...