classifier-toolkit 0.2.7__tar.gz → 0.3.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (209) hide show
  1. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/PKG-INFO +1 -1
  2. classifier_toolkit-0.3.0/classifier_toolkit/explainability/__init__.py +41 -0
  3. classifier_toolkit-0.3.0/classifier_toolkit/explainability/base.py +60 -0
  4. classifier_toolkit-0.3.0/classifier_toolkit/explainability/interactions.py +138 -0
  5. classifier_toolkit-0.3.0/classifier_toolkit/explainability/misclassification.py +107 -0
  6. classifier_toolkit-0.3.0/classifier_toolkit/explainability/plots.py +140 -0
  7. classifier_toolkit-0.3.0/classifier_toolkit/explainability/toolkit.py +137 -0
  8. classifier_toolkit-0.3.0/classifier_toolkit/explainability/tree_explainer.py +136 -0
  9. classifier_toolkit-0.3.0/docs/explainability/interactions.md +48 -0
  10. classifier_toolkit-0.3.0/docs/explainability/misclassification.md +46 -0
  11. classifier_toolkit-0.3.0/docs/explainability/overview.md +87 -0
  12. classifier_toolkit-0.3.0/docs/explainability/plots.md +47 -0
  13. classifier_toolkit-0.3.0/docs/explainability/toolkit.md +59 -0
  14. classifier_toolkit-0.3.0/docs/explainability/tree_explainer.md +50 -0
  15. classifier_toolkit-0.3.0/docs/reference/explainability/interactions.md +7 -0
  16. classifier_toolkit-0.3.0/docs/reference/explainability/misclassification.md +3 -0
  17. classifier_toolkit-0.3.0/docs/reference/explainability/overview.md +42 -0
  18. classifier_toolkit-0.3.0/docs/reference/explainability/plots.md +7 -0
  19. classifier_toolkit-0.3.0/docs/reference/explainability/toolkit.md +3 -0
  20. classifier_toolkit-0.3.0/docs/reference/explainability/tree_explainer.md +7 -0
  21. classifier_toolkit-0.3.0/examples/selected_features.json +189 -0
  22. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/mkdocs.yml +14 -5
  23. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/pyproject.toml +1 -1
  24. classifier_toolkit-0.3.0/tests/explainability/__init__.py +0 -0
  25. classifier_toolkit-0.3.0/tests/explainability/test_catboost_e2e.py +120 -0
  26. classifier_toolkit-0.3.0/tests/explainability/test_misclassification.py +63 -0
  27. classifier_toolkit-0.3.0/tests/explainability/test_plots.py +70 -0
  28. classifier_toolkit-0.3.0/tests/explainability/test_smoke.py +23 -0
  29. classifier_toolkit-0.3.0/tests/explainability/test_toolkit.py +58 -0
  30. classifier_toolkit-0.3.0/tests/explainability/test_tree_explainer.py +50 -0
  31. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/uv.lock +1 -1
  32. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.github/pull_request_template/default.md +0 -0
  33. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.github/workflows/checks.yaml +0 -0
  34. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.github/workflows/docs.yml +0 -0
  35. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.github/workflows/master.yaml +0 -0
  36. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.github/workflows/release.yaml +0 -0
  37. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.github/workflows/working-branch.yaml +0 -0
  38. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.gitignore +0 -0
  39. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.python-version +0 -0
  40. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/.sqlfluff +0 -0
  41. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/LICENSE +0 -0
  42. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/Makefile +0 -0
  43. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/README.md +0 -0
  44. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/data_partition/__init__.py +0 -0
  45. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/data_partition/data_preprocess.py +0 -0
  46. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/data_partition/optimize_data.py +0 -0
  47. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/data_partition/split_train_test.py +0 -0
  48. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/__init__.py +0 -0
  49. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/bivariate_analysis.py +0 -0
  50. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/eda_toolkit.py +0 -0
  51. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/feature_engineering.py +0 -0
  52. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/first_glance.py +0 -0
  53. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/univariate_analysis.py +0 -0
  54. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/visualizations.py +0 -0
  55. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/warnings/__init__.py +0 -0
  56. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/warnings/automated_warnings.py +0 -0
  57. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/eda/warnings/default_warnings.py +0 -0
  58. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/__init__.py +0 -0
  59. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/base.py +0 -0
  60. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/correlation.py +0 -0
  61. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/counter_intuitive.py +0 -0
  62. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/drift.py +0 -0
  63. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/expert_rules.py +0 -0
  64. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/low_variance.py +0 -0
  65. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/predictive_power.py +0 -0
  66. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_reduction/reducer.py +0 -0
  67. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/__init__.py +0 -0
  68. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/base.py +0 -0
  69. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/embedded_methods/__init__.py +0 -0
  70. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/embedded_methods/elastic_net.py +0 -0
  71. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/feature_stability.py +0 -0
  72. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/meta_selector.py +0 -0
  73. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/utils/__init__.py +0 -0
  74. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/utils/data_handling.py +0 -0
  75. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/utils/feature_type_constraints.py +0 -0
  76. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/utils/plottings.py +0 -0
  77. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/utils/scoring.py +0 -0
  78. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/__init__.py +0 -0
  79. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/bayesian_search.py +0 -0
  80. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/boruta.py +0 -0
  81. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/combination_search.py +0 -0
  82. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/recursive_feature_eliminator.py +0 -0
  83. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/rfe.py +0 -0
  84. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/rfe_catboost.py +0 -0
  85. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/feature_selection/wrapper_methods/sequential_selection.py +0 -0
  86. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/__init__.py +0 -0
  87. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/hyper_parameter_tuning/__init__.py +0 -0
  88. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/hyper_parameter_tuning/tuner.py +0 -0
  89. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/models/__init__.py +0 -0
  90. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/models/base.py +0 -0
  91. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/models/ensemble_methods.py +0 -0
  92. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/utils/__init__.py +0 -0
  93. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/classifier_toolkit/model_training/utils/params.py +0 -0
  94. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/changelog.md +0 -0
  95. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/data_partition/data_preprocess.md +0 -0
  96. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/data_partition/optimize_data.md +0 -0
  97. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/data_partition/overview.md +0 -0
  98. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/data_partition/split_train_test.md +0 -0
  99. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/bivariate_analysis.md +0 -0
  100. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/eda_toolkit.md +0 -0
  101. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/feature_engineering.md +0 -0
  102. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/first_glance.md +0 -0
  103. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/overview.md +0 -0
  104. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/univariate_analysis.md +0 -0
  105. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/visualizations.md +0 -0
  106. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/warnings/default_warnings.md +0 -0
  107. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/eda/warnings/warning_system.md +0 -0
  108. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/examples/eda_example.md +0 -0
  109. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/examples/feature_selection_advanced.md +0 -0
  110. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/examples/feature_selection_example.md +0 -0
  111. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/correlation.md +0 -0
  112. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/counter_intuitive.md +0 -0
  113. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/drift.md +0 -0
  114. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/expert_rules.md +0 -0
  115. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/low_variance.md +0 -0
  116. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/overview.md +0 -0
  117. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/predictive_power.md +0 -0
  118. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_reduction/reducer.md +0 -0
  119. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/embedded_methods/elastic_net.md +0 -0
  120. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/feature_stability.md +0 -0
  121. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/meta_selector.md +0 -0
  122. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/overview.md +0 -0
  123. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/utils/data_handling.md +0 -0
  124. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/utils/scoring.md +0 -0
  125. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/wrapper_methods/bayesian_search.md +0 -0
  126. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/wrapper_methods/boruta.md +0 -0
  127. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/wrapper_methods/combination_search.md +0 -0
  128. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/wrapper_methods/recursive_feature_eliminator.md +0 -0
  129. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/wrapper_methods/rfe.md +0 -0
  130. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/feature_selection/wrapper_methods/sequential_selection.md +0 -0
  131. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/index.md +0 -0
  132. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/model_training/tuner.md +0 -0
  133. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/data_partition/data_preprocess.md +0 -0
  134. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/data_partition/optimize_data.md +0 -0
  135. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/data_partition/overview.md +0 -0
  136. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/data_partition/split_train_test.md +0 -0
  137. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/bivariate_analysis.md +0 -0
  138. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/eda_toolkit.md +0 -0
  139. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/feature_engineering.md +0 -0
  140. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/first_glance.md +0 -0
  141. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/overview.md +0 -0
  142. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/univariate_analysis.md +0 -0
  143. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/visualizations.md +0 -0
  144. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/warnings/default_warnings.md +0 -0
  145. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/eda/warnings/warning_system.md +0 -0
  146. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/base.md +0 -0
  147. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/correlation.md +0 -0
  148. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/counter_intuitive.md +0 -0
  149. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/drift.md +0 -0
  150. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/expert_rules.md +0 -0
  151. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/low_variance.md +0 -0
  152. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/overview.md +0 -0
  153. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/predictive_power.md +0 -0
  154. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_reduction/reducer.md +0 -0
  155. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/base.md +0 -0
  156. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/embedded_methods/elastic_net.md +0 -0
  157. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/feature_stability.md +0 -0
  158. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/meta_selector.md +0 -0
  159. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/overview.md +0 -0
  160. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/utils/data_handling.md +0 -0
  161. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/utils/plottings.md +0 -0
  162. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/utils/scoring.md +0 -0
  163. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/wrapper_methods/bayesian_search.md +0 -0
  164. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/wrapper_methods/boruta.md +0 -0
  165. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/wrapper_methods/combination_search.md +0 -0
  166. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/wrapper_methods/recursive_feature_eliminator.md +0 -0
  167. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/wrapper_methods/rfe.md +0 -0
  168. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/feature_selection/wrapper_methods/sequential_selection.md +0 -0
  169. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/model_training/params.md +0 -0
  170. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/reference/tuner/tuner.md +0 -0
  171. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/docs/stylesheets/extra.css +0 -0
  172. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/examples/__init__.py +0 -0
  173. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/examples/example.py +0 -0
  174. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/main.py +0 -0
  175. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/notebooks/paylater_removed.json +0 -0
  176. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/ruff.toml +0 -0
  177. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/__init__.py +0 -0
  178. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/conftest.py +0 -0
  179. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/data_partition/__init__.py +0 -0
  180. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/data_partition/test_data_preprocess.py +0 -0
  181. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/data_partition/test_optimize_data.py +0 -0
  182. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/data_partition/test_split_train_test.py +0 -0
  183. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/__init__.py +0 -0
  184. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/test_bivariate_analysis.py +0 -0
  185. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/test_feature_engineering.py +0 -0
  186. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/test_first_glance.py +0 -0
  187. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/test_univariate_analysis.py +0 -0
  188. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/test_visualizations.py +0 -0
  189. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/eda/test_warnings.py +0 -0
  190. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/__init__.py +0 -0
  191. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_correlation.py +0 -0
  192. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_counter_intuitive.py +0 -0
  193. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_drift.py +0 -0
  194. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_expert_rules.py +0 -0
  195. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_low_variance.py +0 -0
  196. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_predictive_power.py +0 -0
  197. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_reducer.py +0 -0
  198. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_reduction/test_smoke.py +0 -0
  199. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/__init__.py +0 -0
  200. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_bayesian_search.py +0 -0
  201. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_boruta.py +0 -0
  202. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_combination_search.py +0 -0
  203. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_elastic_net.py +0 -0
  204. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_feature_stability.py +0 -0
  205. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_recursive_feature_eliminator.py +0 -0
  206. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_rfe.py +0 -0
  207. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_rfe_catboost.py +0 -0
  208. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_scoring.py +0 -0
  209. {classifier_toolkit-0.2.7 → classifier_toolkit-0.3.0}/tests/feature_selection/test_sequential_selection.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: classifier-toolkit
3
- Version: 0.2.7
3
+ Version: 0.3.0
4
4
  Project-URL: Documentation, http://classifier-toolkit.gh-pages.qonto.co/
5
5
  Author-email: "senih.yilmaz" <senih.yilmaz@qonto.com>, "jeremy.fraoua" <jeremy.fraoua@qonto.com>, "gauthier.marquand" <gauthier.marquand@qonto.com>, "arnaud.alepee" <arnaud.alepee@qonto.com>
6
6
  License-File: LICENSE
@@ -0,0 +1,41 @@
1
+ """Model explainability public API with lazy imports.
2
+
3
+ Exposes SHAP-based global, local, and interaction explanations for tree
4
+ models, mirroring the paylater production explainability workflow.
5
+ """
6
+
7
+ import logging
8
+ from importlib import import_module
9
+ from typing import Any
10
+
11
+ logging.getLogger("classifier_toolkit.explainability").addHandler(logging.NullHandler())
12
+
13
+ _EXPORTS: dict[str, str] = {
14
+ "ExplainabilityError": "classifier_toolkit.explainability.base",
15
+ "ShapResult": "classifier_toolkit.explainability.base",
16
+ "TreeSHAPExplainer": "classifier_toolkit.explainability.tree_explainer",
17
+ "plot_summary": "classifier_toolkit.explainability.plots",
18
+ "plot_waterfall": "classifier_toolkit.explainability.plots",
19
+ "plot_shap_dependence_grid": "classifier_toolkit.explainability.plots",
20
+ "compute_shap_interaction_values": "classifier_toolkit.explainability.interactions",
21
+ "analyze_shap_interactions": "classifier_toolkit.explainability.interactions",
22
+ "plot_top_shap_interactions": "classifier_toolkit.explainability.interactions",
23
+ "MisclassificationAnalyzer": "classifier_toolkit.explainability.misclassification",
24
+ "ExplainabilityToolkit": "classifier_toolkit.explainability.toolkit",
25
+ }
26
+
27
+ __all__ = list(_EXPORTS.keys())
28
+
29
+
30
+ def __getattr__(name: str) -> Any: # pragma: no cover
31
+ module_path = _EXPORTS.get(name)
32
+ if module_path is None:
33
+ raise AttributeError(
34
+ f"module 'classifier_toolkit.explainability' has no attribute '{name}'"
35
+ )
36
+ module = import_module(module_path)
37
+ return getattr(module, name)
38
+
39
+
40
+ def __dir__(): # pragma: no cover
41
+ return sorted(list(globals().keys()) + __all__)
@@ -0,0 +1,60 @@
1
+ """Base types and exceptions for model explainability."""
2
+
3
+ from dataclasses import dataclass
4
+ from typing import Any, Optional, Union
5
+
6
+ import numpy as np
7
+ import pandas as pd
8
+
9
+
10
+ class ExplainabilityError(Exception):
11
+ """Base exception for explainability errors."""
12
+
13
+
14
+ @dataclass
15
+ class ShapResult:
16
+ """Container for SHAP explanation output.
17
+
18
+ Attributes
19
+ ----------
20
+ values : np.ndarray
21
+ SHAP values with shape ``(n_samples, n_features)`` for binary
22
+ classification, or ``(n_samples, n_features, n_classes)`` for
23
+ multiclass models.
24
+ data : Union[pd.DataFrame, np.ndarray]
25
+ Feature matrix that was explained.
26
+ feature_names : List[str]
27
+ Names of the explained features.
28
+ base_values : Optional[np.ndarray]
29
+ Expected model output (SHAP base value) per sample.
30
+ explainer : Any
31
+ Underlying SHAP explainer instance (for waterfall / interaction reuse).
32
+ """
33
+
34
+ values: np.ndarray
35
+ data: Union[pd.DataFrame, np.ndarray]
36
+ feature_names: list[str]
37
+ base_values: Optional[np.ndarray] = None
38
+ explainer: Any = None
39
+
40
+ @property
41
+ def n_samples(self) -> int:
42
+ return int(self.values.shape[0])
43
+
44
+ @property
45
+ def n_features(self) -> int:
46
+ if self.values.ndim == 2:
47
+ return int(self.values.shape[1])
48
+ return int(self.values.shape[1])
49
+
50
+ def values_for_class(self, class_index: int = -1) -> np.ndarray:
51
+ """Return a 2D SHAP matrix for plotting (positive class by default)."""
52
+ if self.values.ndim == 3:
53
+ return self.values[..., class_index]
54
+ return self.values
55
+
56
+ def mean_abs_importance(self) -> pd.Series:
57
+ """Mean absolute SHAP value per feature (global importance)."""
58
+ values = self.values_for_class()
59
+ importance = np.abs(values).mean(axis=0)
60
+ return pd.Series(importance, index=self.feature_names, name="mean_abs_shap")
@@ -0,0 +1,138 @@
1
+ """SHAP interaction analysis and plotting."""
2
+
3
+ import math
4
+ from collections.abc import Sequence
5
+ from typing import Union
6
+
7
+ import matplotlib.pyplot as plt
8
+ import numpy as np
9
+ import pandas as pd
10
+ import shap
11
+
12
+
13
+ def _normalize_shap_array(values, class_index: int = -1) -> np.ndarray:
14
+ """Convert list-shaped multiclass SHAP output to a single 2D/3D array."""
15
+ if isinstance(values, list):
16
+ return np.asarray(values[class_index])
17
+ return np.asarray(values)
18
+
19
+
20
+ def compute_shap_interaction_values(
21
+ model, X: pd.DataFrame, feature_names: Sequence[str], class_index: int = -1
22
+ ) -> np.ndarray:
23
+ """Compute SHAP interaction values for tree models."""
24
+ raw = shap.TreeExplainer(model).shap_interaction_values(X[feature_names])
25
+ return _normalize_shap_array(raw, class_index=class_index)
26
+
27
+
28
+ def analyze_shap_interactions(
29
+ model,
30
+ X_test: pd.DataFrame,
31
+ feature_names: Sequence[str],
32
+ shap_values: np.ndarray,
33
+ n_important_features: int = 2,
34
+ ) -> plt.Figure:
35
+ """Grid of main effects (diagonal) and pairwise interactions (off-diagonal)."""
36
+ interaction_values = compute_shap_interaction_values(model, X_test, feature_names)
37
+
38
+ important_features = list(feature_names)[:n_important_features]
39
+ n_features = len(important_features)
40
+
41
+ fig, axes = plt.subplots(
42
+ n_features, n_features, figsize=(10 * n_features, 10 * n_features)
43
+ )
44
+
45
+ if n_features == 1:
46
+ axes = np.array([[axes]])
47
+
48
+ feature_list = list(feature_names)
49
+ for i in range(n_features):
50
+ for j in range(n_features):
51
+ feature_i = important_features[i]
52
+ feature_j = important_features[j]
53
+ idx_i = feature_list.index(feature_i)
54
+ idx_j = feature_list.index(feature_j)
55
+ plt.sca(axes[i, j])
56
+
57
+ if i == j:
58
+ shap.dependence_plot(
59
+ ind=idx_i,
60
+ shap_values=shap_values,
61
+ features=X_test[feature_list],
62
+ feature_names=feature_list,
63
+ ax=axes[i, j],
64
+ show=False,
65
+ )
66
+ axes[i, j].set_title(f"Main effect: {feature_i}")
67
+ else:
68
+ shap.dependence_plot(
69
+ (idx_i, idx_j),
70
+ interaction_values,
71
+ X_test[feature_list],
72
+ feature_names=feature_list,
73
+ ax=axes[i, j],
74
+ show=False,
75
+ )
76
+ axes[i, j].set_title(f"Interaction: {feature_i} × {feature_j}") # noqa: RUF001
77
+
78
+ plt.tight_layout()
79
+ plt.subplots_adjust(top=0.9)
80
+ fig.suptitle("Main Effects and Interactions Analysis", fontsize=16)
81
+ return fig
82
+
83
+
84
+ def plot_top_shap_interactions(
85
+ model,
86
+ X_test: Union[pd.DataFrame, np.ndarray],
87
+ feature_names: Sequence[str],
88
+ n_top_interactions: int = 9,
89
+ ) -> plt.Figure:
90
+ """Plot the strongest feature interactions by mean absolute SHAP interaction."""
91
+ if not isinstance(X_test, pd.DataFrame):
92
+ X_test = pd.DataFrame(X_test, columns=list(feature_names))
93
+
94
+ interaction_values = compute_shap_interaction_values(model, X_test, feature_names)
95
+
96
+ interaction_strength = np.abs(interaction_values).mean(axis=0)
97
+ np.fill_diagonal(interaction_strength, 0)
98
+
99
+ flat_indices = np.argsort(interaction_strength.flatten())[-n_top_interactions:]
100
+ feature_pairs: list[tuple] = []
101
+ names = list(feature_names)
102
+
103
+ for flat_idx in flat_indices[::-1]:
104
+ i, j = np.unravel_index(flat_idx, interaction_strength.shape)
105
+ if i < j:
106
+ feature_pairs.append((names[i], names[j], interaction_strength[i, j]))
107
+ else:
108
+ feature_pairs.append((names[j], names[i], interaction_strength[i, j]))
109
+
110
+ n_cols = 2
111
+ n_rows = math.ceil(len(feature_pairs) / n_cols)
112
+ fig, axes = plt.subplots(n_rows, n_cols, figsize=(18, 6 * n_rows))
113
+
114
+ if n_rows == 1:
115
+ axes = np.array([axes])
116
+ axes_flat = axes.flatten()
117
+
118
+ for i, (feature1, feature2, strength) in enumerate(feature_pairs):
119
+ ax = axes_flat[i]
120
+ idx1 = names.index(feature1)
121
+ idx2 = names.index(feature2)
122
+ plt.sca(ax)
123
+ shap.dependence_plot(
124
+ (idx1, idx2),
125
+ interaction_values,
126
+ X_test[names],
127
+ feature_names=names,
128
+ ax=ax,
129
+ show=False,
130
+ )
131
+ ax.set_title(f"#{i + 1}: {feature1} × {feature2}\nStrength: {strength:.4f}") # noqa: RUF001
132
+
133
+ for i in range(len(feature_pairs), n_rows * n_cols):
134
+ axes_flat[i].set_visible(False)
135
+
136
+ fig.suptitle("Top Feature Interactions by SHAP Strength", fontsize=16, y=1.02)
137
+ plt.tight_layout()
138
+ return fig
@@ -0,0 +1,107 @@
1
+ """Misclassification analysis with per-observation SHAP waterfalls."""
2
+
3
+ from collections.abc import Sequence
4
+ from typing import Literal, Optional, Union
5
+
6
+ import numpy as np
7
+ import pandas as pd
8
+
9
+ from classifier_toolkit.explainability.base import ExplainabilityError
10
+ from classifier_toolkit.explainability.plots import plot_waterfall
11
+ from classifier_toolkit.explainability.tree_explainer import TreeSHAPExplainer
12
+
13
+ Quadrant = Literal[
14
+ "false_negative",
15
+ "false_positive",
16
+ "true_positive",
17
+ "true_negative",
18
+ ]
19
+
20
+ _QUADRANT_FILTERS: dict[Quadrant, tuple] = {
21
+ "false_negative": ("y_pred", False, "y_true", True),
22
+ "false_positive": ("y_pred", True, "y_true", False),
23
+ "true_positive": ("y_pred", True, "y_true", True),
24
+ "true_negative": ("y_pred", False, "y_true", False),
25
+ }
26
+
27
+ _QUADRANT_SORT: dict[Quadrant, tuple] = {
28
+ "false_negative": ("y_pred_proba", True),
29
+ "false_positive": ("y_pred_proba", False),
30
+ "true_positive": ("y_pred_proba", False),
31
+ "true_negative": ("y_pred_proba", True),
32
+ }
33
+
34
+
35
+ class MisclassificationAnalyzer:
36
+ """Filter confusion-matrix quadrants and explain individual predictions.
37
+
38
+ Mirrors the bin-analysis workflow from the paylater explainability
39
+ notebook (false negatives, false positives, true positives, true negatives).
40
+ """
41
+
42
+ def __init__(
43
+ self,
44
+ model,
45
+ explainer: TreeSHAPExplainer,
46
+ feature_names: Sequence[str],
47
+ threshold: float = 0.5,
48
+ ) -> None:
49
+ self.model = model
50
+ self.explainer = explainer
51
+ self.feature_names = list(feature_names)
52
+ self.threshold = threshold
53
+
54
+ def build_prediction_frame(
55
+ self,
56
+ X: pd.DataFrame,
57
+ y_true: Union[pd.Series, np.ndarray],
58
+ extra_columns: Optional[pd.DataFrame] = None,
59
+ ) -> pd.DataFrame:
60
+ """Attach true labels, predicted probabilities, and binary predictions."""
61
+ df = X[self.feature_names].copy()
62
+ if extra_columns is not None:
63
+ for col in extra_columns.columns:
64
+ df[col] = extra_columns[col].values
65
+ df["y_true"] = np.asarray(y_true)
66
+ proba = self.model.predict_proba(df[self.feature_names])[:, 1]
67
+ df["y_pred_proba"] = proba
68
+ df["y_pred"] = (proba > self.threshold).astype(int)
69
+ return df
70
+
71
+ def filter_quadrant(
72
+ self,
73
+ df: pd.DataFrame,
74
+ quadrant: Quadrant,
75
+ ) -> pd.DataFrame:
76
+ """Return rows for a confusion-matrix quadrant, sorted by probability."""
77
+ if quadrant not in _QUADRANT_FILTERS:
78
+ raise ExplainabilityError(f"Unknown quadrant: {quadrant}")
79
+
80
+ pred_col, pred_val, true_col, true_val = _QUADRANT_FILTERS[quadrant]
81
+ mask = (df[pred_col] == pred_val) & (df[true_col] == true_val)
82
+ sort_col, ascending = _QUADRANT_SORT[quadrant]
83
+ return df.loc[mask].sort_values(by=sort_col, ascending=ascending)
84
+
85
+ def explain_row(self, row: Union[pd.Series, pd.DataFrame]):
86
+ """SHAP explanation object for one observation (waterfall-ready)."""
87
+ if isinstance(row, pd.Series):
88
+ row = row.to_frame().T
89
+ if len(row) != 1:
90
+ raise ExplainabilityError("explain_row expects a single observation.")
91
+ return self.explainer.explain_observation(row[self.feature_names])
92
+
93
+ def plot_waterfall_for_row(
94
+ self,
95
+ row: Union[pd.Series, pd.DataFrame],
96
+ max_display: int = 15,
97
+ show: bool = True,
98
+ ) -> None:
99
+ """Waterfall plot for one observation."""
100
+ plot_waterfall(self.explain_row(row), max_display=max_display, show=show)
101
+
102
+ def quadrant_counts(self, df: pd.DataFrame) -> pd.Series:
103
+ """Count rows in each confusion-matrix quadrant."""
104
+ counts = {}
105
+ for quadrant in _QUADRANT_FILTERS:
106
+ counts[quadrant] = len(self.filter_quadrant(df, quadrant))
107
+ return pd.Series(counts, name="count")
@@ -0,0 +1,140 @@
1
+ """SHAP plotting helpers."""
2
+
3
+ import math
4
+ from collections.abc import Sequence
5
+ from typing import Optional, Union
6
+
7
+ import matplotlib.pyplot as plt
8
+ import numpy as np
9
+ import pandas as pd
10
+ import shap
11
+
12
+ from classifier_toolkit.explainability.base import ShapResult
13
+
14
+
15
+ def plot_summary(
16
+ shap_result: ShapResult,
17
+ max_display: int = 20,
18
+ show: bool = True,
19
+ ) -> None:
20
+ """Beeswarm summary plot for global feature importance."""
21
+ shap.summary_plot(
22
+ shap_result.values_for_class(),
23
+ shap_result.data,
24
+ feature_names=shap_result.feature_names,
25
+ max_display=max_display,
26
+ show=show,
27
+ )
28
+
29
+
30
+ def plot_waterfall(
31
+ explanation,
32
+ max_display: int = 15,
33
+ show: bool = True,
34
+ ) -> None:
35
+ """Waterfall plot for a single observation SHAP explanation."""
36
+ shap.plots.waterfall(explanation, max_display=max_display, show=show)
37
+
38
+
39
+ def plot_shap_dependence_grid(
40
+ feature_list: Union[str, Sequence[str]],
41
+ shap_values: Union[np.ndarray, pd.DataFrame],
42
+ X_test: pd.DataFrame,
43
+ feature_names: Optional[Sequence[str]] = None,
44
+ n_cols: int = 2,
45
+ overall_title: str = "Feature Impact Analysis using SHAP",
46
+ is_categorical: bool = False,
47
+ show: bool = False,
48
+ ) -> Optional[plt.Figure]:
49
+ """Plot SHAP dependence plots for a list of features in a grid layout.
50
+
51
+ Ported from the paylater production workflow to provide the same
52
+ grid-layout dependence analysis inside classifier_toolkit.
53
+ """
54
+ if isinstance(feature_list, str):
55
+ feature_list = [feature_list]
56
+ feature_list = list(feature_list)
57
+
58
+ if len(feature_list) == 0:
59
+ return None
60
+
61
+ if feature_names is None:
62
+ feature_names = list(X_test.columns)
63
+ feature_names = list(feature_names)
64
+
65
+ missing = [f for f in feature_list if f not in feature_names]
66
+ if missing:
67
+ raise ValueError(
68
+ f"The following features are not in feature_names: {missing}. "
69
+ f"Available features: {feature_names}"
70
+ )
71
+
72
+ n_features = len(feature_list)
73
+ n_rows = math.ceil(n_features / n_cols)
74
+
75
+ fig, axes = plt.subplots(n_rows, n_cols, figsize=(18, 6 * n_rows))
76
+
77
+ if n_rows * n_cols == 1:
78
+ axes_flat = [axes]
79
+ elif n_rows == 1:
80
+ axes_flat = list(axes)
81
+ else:
82
+ axes_flat = list(axes.flatten())
83
+
84
+ if not is_categorical:
85
+ for i, feature in enumerate(feature_list):
86
+ ax = axes_flat[i]
87
+ try:
88
+ shap.dependence_plot(
89
+ ind=feature,
90
+ shap_values=shap_values,
91
+ features=X_test[feature_names],
92
+ feature_names=feature_names,
93
+ ax=ax,
94
+ show=False,
95
+ )
96
+ ax.set_title(f"SHAP Dependence: {feature}")
97
+ except Exception as exc: # pragma: no cover - plotting edge cases
98
+ ax.text(
99
+ 0.5,
100
+ 0.5,
101
+ f"Error plotting\n{feature}\n{exc}",
102
+ ha="center",
103
+ va="center",
104
+ transform=ax.transAxes,
105
+ )
106
+ else:
107
+ for i, feature in enumerate(feature_list):
108
+ row = i // n_cols
109
+ col = i % n_cols
110
+ feature_idx = feature_names.index(feature)
111
+ ax = axes[row, col] if n_rows > 1 else axes_flat[i]
112
+ try:
113
+ shap.dependence_plot(
114
+ feature_idx,
115
+ shap_values,
116
+ X_test,
117
+ title=f"Dependence: {feature}",
118
+ ax=ax,
119
+ show=False,
120
+ interaction_index=None,
121
+ )
122
+ except Exception as exc: # pragma: no cover
123
+ ax.text(
124
+ 0.5,
125
+ 0.5,
126
+ f"Error plotting\n{feature}\n{exc}",
127
+ ha="center",
128
+ va="center",
129
+ transform=ax.transAxes,
130
+ )
131
+
132
+ if n_rows * n_cols > 1:
133
+ for i in range(n_features, n_rows * n_cols):
134
+ axes_flat[i].set_visible(False)
135
+
136
+ fig.suptitle(overall_title, fontsize=16, y=1.02)
137
+ plt.tight_layout()
138
+ if show:
139
+ plt.show()
140
+ return fig
@@ -0,0 +1,137 @@
1
+ """High-level explainability orchestrator."""
2
+
3
+ from collections.abc import Sequence
4
+ from typing import Optional, Union
5
+
6
+ import numpy as np
7
+ import pandas as pd
8
+
9
+ from classifier_toolkit.explainability.base import ShapResult
10
+ from classifier_toolkit.explainability.interactions import (
11
+ analyze_shap_interactions,
12
+ plot_top_shap_interactions,
13
+ )
14
+ from classifier_toolkit.explainability.misclassification import (
15
+ MisclassificationAnalyzer,
16
+ )
17
+ from classifier_toolkit.explainability.plots import (
18
+ plot_shap_dependence_grid,
19
+ plot_summary,
20
+ )
21
+ from classifier_toolkit.explainability.tree_explainer import TreeSHAPExplainer
22
+
23
+
24
+ class ExplainabilityToolkit:
25
+ """Orchestrate the standard SHAP explainability workflow.
26
+
27
+ Combines global summaries, dependence grids, interaction analysis, and
28
+ misclassification waterfalls behind a single entry point.
29
+ """
30
+
31
+ def __init__(
32
+ self,
33
+ model,
34
+ feature_names: Sequence[str],
35
+ numerical_features: Optional[Sequence[str]] = None,
36
+ categorical_features: Optional[Sequence[str]] = None,
37
+ cat_features: Optional[list[str]] = None,
38
+ threshold: float = 0.5,
39
+ ) -> None:
40
+ self.model = model
41
+ self.feature_names = list(feature_names)
42
+ self.numerical_features = list(numerical_features or [])
43
+ self.categorical_features = list(categorical_features or [])
44
+ self.explainer = TreeSHAPExplainer(
45
+ model,
46
+ cat_features=cat_features or categorical_features,
47
+ feature_names=self.feature_names,
48
+ )
49
+ self.threshold = threshold
50
+ self.shap_result_: Optional[ShapResult] = None
51
+
52
+ def compute_shap(
53
+ self,
54
+ X: pd.DataFrame,
55
+ y: Optional[Union[pd.Series, np.ndarray]] = None,
56
+ ) -> ShapResult:
57
+ """Compute and cache SHAP values for a feature matrix."""
58
+ self.shap_result_ = self.explainer.explain(X[self.feature_names], y=y)
59
+ return self.shap_result_
60
+
61
+ def plot_global_summary(self, max_display: int = 20, show: bool = True) -> None:
62
+ """Beeswarm summary plot from cached SHAP values."""
63
+ if self.shap_result_ is None:
64
+ raise ValueError("Call compute_shap() first.")
65
+ plot_summary(self.shap_result_, max_display=max_display, show=show)
66
+
67
+ def plot_numerical_dependence(self, show: bool = False):
68
+ """Dependence grid for numerical features."""
69
+ if self.shap_result_ is None:
70
+ raise ValueError("Call compute_shap() first.")
71
+ if not self.numerical_features:
72
+ return None
73
+ n_num = len(self.numerical_features)
74
+ values_2d = self.shap_result_.values_for_class()
75
+ shap_subset = values_2d[:, :n_num]
76
+ X_num = self.shap_result_.data[self.numerical_features]
77
+ return plot_shap_dependence_grid(
78
+ feature_list=self.numerical_features,
79
+ shap_values=shap_subset,
80
+ X_test=X_num,
81
+ feature_names=self.numerical_features,
82
+ overall_title="Numerical Feature Impact (SHAP)",
83
+ show=show,
84
+ )
85
+
86
+ def plot_categorical_dependence(self, show: bool = False):
87
+ """Dependence grid for categorical features."""
88
+ if self.shap_result_ is None:
89
+ raise ValueError("Call compute_shap() first.")
90
+ if not self.categorical_features:
91
+ return None
92
+ n_num = len(self.numerical_features)
93
+ values_2d = self.shap_result_.values_for_class()
94
+ shap_cat = values_2d[:, n_num:]
95
+ X_cat = self.shap_result_.data[self.categorical_features]
96
+ return plot_shap_dependence_grid(
97
+ feature_list=self.categorical_features,
98
+ shap_values=shap_cat,
99
+ X_test=X_cat,
100
+ feature_names=self.categorical_features,
101
+ is_categorical=True,
102
+ overall_title="Categorical Feature Impact (SHAP)",
103
+ show=show,
104
+ )
105
+
106
+ def plot_interaction_grid(self, n_important_features: int = 2):
107
+ """Main-effect / interaction grid for top features."""
108
+ if self.shap_result_ is None:
109
+ raise ValueError("Call compute_shap() first.")
110
+ X = self.shap_result_.data
111
+ return analyze_shap_interactions(
112
+ self.model,
113
+ X,
114
+ self.feature_names,
115
+ self.shap_result_.values_for_class(),
116
+ n_important_features=n_important_features,
117
+ )
118
+
119
+ def plot_top_interactions(self, n_top_interactions: int = 5):
120
+ """Top pairwise SHAP interactions."""
121
+ if self.shap_result_ is None:
122
+ raise ValueError("Call compute_shap() first.")
123
+ return plot_top_shap_interactions(
124
+ self.model,
125
+ self.shap_result_.data,
126
+ self.feature_names,
127
+ n_top_interactions=n_top_interactions,
128
+ )
129
+
130
+ def misclassification_analyzer(self) -> MisclassificationAnalyzer:
131
+ """Analyzer for confusion-matrix quadrant waterfalls."""
132
+ return MisclassificationAnalyzer(
133
+ model=self.model,
134
+ explainer=self.explainer,
135
+ feature_names=self.feature_names,
136
+ threshold=self.threshold,
137
+ )