classifier-toolkit 0.4.0__tar.gz → 0.5.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 (156) hide show
  1. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/PKG-INFO +2 -2
  2. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/data_partition/split_train_test.py +75 -154
  3. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/interactions.py +39 -19
  4. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/toolkit.py +17 -4
  5. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/correlation.py +234 -43
  6. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/reducer.py +18 -2
  7. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/utils/scoring.py +0 -33
  8. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/combination_search.py +241 -176
  9. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/recursive_feature_eliminator.py +467 -23
  10. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/rfe.py +0 -32
  11. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/hyper_parameter_tuning/tuner.py +211 -37
  12. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/utils/params.py +33 -9
  13. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_validation/__init__.py +16 -0
  14. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_validation/evaluator.py +593 -10
  15. classifier_toolkit-0.5.0/classifier_toolkit/model_validation/model_assessment.py +975 -0
  16. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_combination_feature_search.ipynb +159 -26
  17. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_explainability_catboost.ipynb +4 -8
  18. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_explainability_lgbm.ipynb +4 -8
  19. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_feature_reduction.ipynb +4 -6
  20. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_train_test_partition.ipynb +2 -2
  21. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/pyproject.toml +2 -2
  22. classifier_toolkit-0.5.0/tests/data_partition/test_split_train_test.py +247 -0
  23. classifier_toolkit-0.5.0/tests/explainability/test_interactions_ranking.py +134 -0
  24. classifier_toolkit-0.5.0/tests/feature_reduction/test_correlation_tie_break.py +262 -0
  25. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_reducer.py +68 -0
  26. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_combination_search.py +311 -2
  27. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_recursive_feature_eliminator.py +192 -0
  28. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_rfe.py +6 -17
  29. classifier_toolkit-0.5.0/tests/feature_selection/test_rfe_assessment.py +439 -0
  30. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_scoring.py +4 -13
  31. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_training/hyper_parameter_tuning/test_params.py +62 -0
  32. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_training/hyper_parameter_tuning/test_tuner.py +451 -0
  33. classifier_toolkit-0.5.0/tests/model_validation/test_evaluator_multiclass.py +313 -0
  34. classifier_toolkit-0.5.0/tests/model_validation/test_model_assessment.py +760 -0
  35. classifier_toolkit-0.4.0/tests/data_partition/test_split_train_test.py +0 -122
  36. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/.gitignore +0 -0
  37. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/LICENSE +0 -0
  38. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/README.md +0 -0
  39. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/calibration/__init__.py +0 -0
  40. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/calibration/base.py +0 -0
  41. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/calibration/ovr_calibration.py +0 -0
  42. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/calibration/reliability.py +0 -0
  43. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/data_partition/__init__.py +0 -0
  44. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/data_partition/data_preprocess.py +0 -0
  45. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/data_partition/optimize_data.py +0 -0
  46. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/datasets/__init__.py +0 -0
  47. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/datasets/_demo.py +0 -0
  48. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/__init__.py +0 -0
  49. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/bivariate_analysis.py +0 -0
  50. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/eda_toolkit.py +0 -0
  51. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/feature_engineering.py +0 -0
  52. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/first_glance.py +0 -0
  53. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/univariate_analysis.py +0 -0
  54. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/visualizations.py +0 -0
  55. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/warnings/__init__.py +0 -0
  56. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/warnings/automated_warnings.py +0 -0
  57. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/eda/warnings/default_warnings.py +0 -0
  58. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/__init__.py +0 -0
  59. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/base.py +0 -0
  60. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/misclassification.py +0 -0
  61. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/plots.py +0 -0
  62. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/explainability/tree_explainer.py +0 -0
  63. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/__init__.py +0 -0
  64. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/base.py +0 -0
  65. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/counter_intuitive.py +0 -0
  66. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/drift.py +0 -0
  67. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/expert_rules.py +0 -0
  68. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/low_variance.py +0 -0
  69. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_reduction/predictive_power.py +0 -0
  70. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/__init__.py +0 -0
  71. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/base.py +0 -0
  72. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/embedded_methods/__init__.py +0 -0
  73. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/embedded_methods/elastic_net.py +0 -0
  74. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/feature_stability.py +0 -0
  75. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/meta_selector.py +0 -0
  76. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/utils/__init__.py +0 -0
  77. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/utils/data_handling.py +0 -0
  78. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/utils/feature_type_constraints.py +0 -0
  79. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/utils/plottings.py +0 -0
  80. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/__init__.py +0 -0
  81. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/bayesian_search.py +0 -0
  82. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/boruta.py +0 -0
  83. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/rfe_catboost.py +0 -0
  84. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/feature_selection/wrapper_methods/sequential_selection.py +0 -0
  85. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/__init__.py +0 -0
  86. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/hyper_parameter_tuning/__init__.py +0 -0
  87. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/models/__init__.py +0 -0
  88. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/models/base.py +0 -0
  89. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/models/ensemble_methods.py +0 -0
  90. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/model_training/utils/__init__.py +0 -0
  91. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/risk_class/__init__.py +0 -0
  92. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/risk_class/risk_classes.py +0 -0
  93. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/classifier_toolkit/risk_class/risk_classes_dp.py +0 -0
  94. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_bayesian_search.ipynb +0 -0
  95. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_grid_search.ipynb +0 -0
  96. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_model_training_catboost.ipynb +0 -0
  97. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_model_training_lgbm.ipynb +0 -0
  98. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_model_validation.ipynb +0 -0
  99. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_recursive_feature_eliminator.ipynb +0 -0
  100. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/examples/example_risk_classes_dp.ipynb +0 -0
  101. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/__init__.py +0 -0
  102. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/calibration/__init__.py +0 -0
  103. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/calibration/test_ovr_calibration.py +0 -0
  104. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/calibration/test_reliability.py +0 -0
  105. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/conftest.py +0 -0
  106. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/data_partition/__init__.py +0 -0
  107. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/data_partition/test_data_preprocess.py +0 -0
  108. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/data_partition/test_optimize_data.py +0 -0
  109. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/datasets/__init__.py +0 -0
  110. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/datasets/test_demo_data.py +0 -0
  111. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/__init__.py +0 -0
  112. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/test_bivariate_analysis.py +0 -0
  113. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/test_feature_engineering.py +0 -0
  114. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/test_first_glance.py +0 -0
  115. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/test_univariate_analysis.py +0 -0
  116. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/test_visualizations.py +0 -0
  117. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/eda/test_warnings.py +0 -0
  118. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/__init__.py +0 -0
  119. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_catboost_e2e.py +0 -0
  120. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_interactions.py +0 -0
  121. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_misclassification.py +0 -0
  122. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_multiclass_e2e.py +0 -0
  123. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_plots.py +0 -0
  124. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_smoke.py +0 -0
  125. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_toolkit.py +0 -0
  126. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/explainability/test_tree_explainer.py +0 -0
  127. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/__init__.py +0 -0
  128. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_correlation.py +0 -0
  129. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_counter_intuitive.py +0 -0
  130. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_drift.py +0 -0
  131. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_expert_rules.py +0 -0
  132. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_low_variance.py +0 -0
  133. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_predictive_power.py +0 -0
  134. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_reduction/test_smoke.py +0 -0
  135. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/__init__.py +0 -0
  136. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_bayesian_search.py +0 -0
  137. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_boruta.py +0 -0
  138. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_elastic_net.py +0 -0
  139. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_feature_stability.py +0 -0
  140. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_rfe_catboost.py +0 -0
  141. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/feature_selection/test_sequential_selection.py +0 -0
  142. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_training/__init__.py +0 -0
  143. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_training/hyper_parameter_tuning/__init__.py +0 -0
  144. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_training/models/__init__.py +0 -0
  145. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_training/models/test_ensemble_methods.py +0 -0
  146. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/__init__.py +0 -0
  147. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/test_compare_score_distributions.py +0 -0
  148. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/test_evaluate_risk_classes_dp.py +0 -0
  149. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/test_evaluator_threshold.py +0 -0
  150. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/test_print_full_validation_report.py +0 -0
  151. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/test_print_risk_class_validation_report.py +0 -0
  152. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/model_validation/test_score_distribution_psi.py +0 -0
  153. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/risk_class/__init__.py +0 -0
  154. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/risk_class/test_construct_bins_dp.py +0 -0
  155. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/risk_class/test_risk_classes.py +0 -0
  156. {classifier_toolkit-0.4.0 → classifier_toolkit-0.5.0}/tests/risk_class/test_validate_risk_classes.py +0 -0
@@ -1,10 +1,10 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: classifier-toolkit
3
- Version: 0.4.0
3
+ Version: 0.5.0
4
4
  Project-URL: Documentation, https://supreme-adventure-jg5qkyr.pages.github.io/
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
7
- Requires-Python: <3.14,>=3.9
7
+ Requires-Python: <3.14,>=3.10
8
8
  Requires-Dist: catboost<2.0.0,>=1.2.2
9
9
  Requires-Dist: category-encoders<3.0.0,>=2.6.3
10
10
  Requires-Dist: colorama<0.5.0,>=0.4.6
@@ -11,200 +11,121 @@ def split_data_recent_orgs(
11
11
  date_col="observation_date",
12
12
  ):
13
13
  """
14
- Split data into train and test sets based on the most recent organizations.
15
- Takes the most recent organizations (by first observation date) for the test set
16
- based on the specified test_size proportion, and the remaining organizations for training.
17
- Iteratively adjusts the split to maintain default rate balance within tolerance.
14
+ Split data into train and test sets on the most recent entities.
15
+
16
+ Entities (``org_id_col``) are ordered by their first observation date, and
17
+ the most recent ones go to the test set, so the test set mimics entities
18
+ the model will meet after deployment. The number of test entities is
19
+ ``test_size`` of all entities, adjusted as little as possible so that the
20
+ train and test event rates differ by at most ``rate_tolerance``: sizes are
21
+ tried from the target outwards (n, n-1, n+1, n-2, ...) within +/-30% of
22
+ it (at least +/-5), and the first one within tolerance is kept. If none
23
+ is, the size with the smallest event-rate difference is used. The target
24
+ is ``round(test_size * n_entities)``, at least 1.
18
25
 
19
26
  Parameters:
20
27
  -----------
21
28
  df : pandas.DataFrame
22
29
  The input dataframe
23
30
  test_size : float, default=0.2
24
- Proportion of organizations to include in test set (most recent ones)
31
+ Share of **entities** (not rows) put in the test set. Recent entities
32
+ usually have fewer observations, so the test share of rows is lower.
25
33
  rate_tolerance : float, default=0.05
26
- Maximum acceptable difference in default rates between train and test (5%)
27
- max_iterations : int, default=50
28
- Maximum number of iterations to find a balanced split
34
+ Maximum absolute difference between the train and test event rates
35
+ (``0.05`` = 5 percentage points; use a value in proportion to the
36
+ event rate).
29
37
  label_col : str, default='default_label'
30
- Column containing the binary labels
38
+ Column containing the binary (0/1) target
31
39
  org_id_col : str, default='org_id'
32
- Column containing organization IDs
40
+ Column containing the entity IDs. Entities that share a first
41
+ observation date are ordered by ID, so the IDs must be mutually
42
+ comparable (not a mix of ``int`` and ``str``).
33
43
  date_col : str, default='observation_date'
34
44
  Column containing the observation dates
35
45
 
36
46
  Returns:
37
47
  --------
38
48
  tuple
39
- (train_df, test_df) - Training and testing dataframes
49
+ ``(train_df, test_df, stats)``, where ``stats`` is the same dict as
50
+ the other splitters return (``train_orgs``, ``test_orgs``,
51
+ ``train_records``, ``test_records``, ``train_target_rate``,
52
+ ``test_target_rate``, ``target_rate_diff``, ...).
40
53
  """
41
- # Ensure date column is datetime
42
54
  df = df.copy()
43
55
  if not pd.api.types.is_datetime64_dtype(df[date_col]):
44
56
  df[date_col] = pd.to_datetime(df[date_col])
45
57
 
46
- # Get the earliest observation date for each organization
58
+ # First observation date of each entity, oldest first. The entity ID
59
+ # breaks ties, so which entities of a same-date group end up in test is
60
+ # deterministic.
47
61
  org_first_dates = (
48
62
  df.groupby(org_id_col)[date_col]
49
63
  .min()
50
64
  .reset_index()
51
- .sort_values(date_col, ascending=True)
65
+ .sort_values([date_col, org_id_col], kind="stable")
52
66
  )
53
-
54
- # Calculate number of organizations for test set
55
67
  total_orgs = len(org_first_dates)
56
- n_test_orgs = int(total_orgs * test_size)
57
-
58
- print("=== RECENT ORGANIZATIONS SPLIT WITH RATE TOLERANCE ===")
59
- print(f"Test size target: {test_size:.1%}")
60
- print(f"Rate tolerance: {rate_tolerance:.1%}")
61
- print(f"Total organizations: {total_orgs}")
68
+ if total_orgs < 2:
69
+ raise ValueError(
70
+ f"split_data_recent_orgs needs at least 2 entities in {org_id_col!r}, "
71
+ f"got {total_orgs}."
72
+ )
73
+ n_test_orgs = max(1, round(total_orgs * test_size))
62
74
 
63
- # Try different splits around the target to find one within tolerance
64
- best_split = None
65
- best_rate_diff = float("inf")
75
+ print("=== RECENT ENTITIES SPLIT WITH RATE TOLERANCE ===")
76
+ print(f"Test size target: {test_size:.1%} of {total_orgs} entities")
77
+ print(f"Rate tolerance: {rate_tolerance:.2%}")
66
78
 
67
- # Search range: ±30% of target test size
79
+ # Sizes from the target outwards, so the first one within tolerance is
80
+ # the closest to the requested test_size.
68
81
  search_range = max(5, int(n_test_orgs * 0.3))
69
- min_test_orgs = max(1, n_test_orgs - search_range)
70
- max_test_orgs = min(total_orgs - 1, n_test_orgs + search_range)
71
-
72
- print(
73
- f"Searching for balanced split (testing {min_test_orgs} to {max_test_orgs} orgs)..."
74
- )
75
-
76
- for test_org_count in range(min_test_orgs, max_test_orgs + 1):
77
- # Take organizations for this test size
78
- test_orgs_data = org_first_dates.tail(test_org_count)
79
-
80
- # Create temporary splits
81
- temp_train = df[~df[org_id_col].isin(test_orgs_data[org_id_col].values)]
82
- temp_test = df[df[org_id_col].isin(test_orgs_data[org_id_col].values)]
83
-
84
- if len(temp_train) == 0 or len(temp_test) == 0:
82
+ low = max(1, n_test_orgs - search_range)
83
+ high = min(total_orgs - 1, n_test_orgs + search_range)
84
+ candidates = sorted(range(low, high + 1), key=lambda n: (abs(n - n_test_orgs), n))
85
+
86
+ best_count, best_rate_diff = None, float("inf")
87
+ for test_org_count in candidates:
88
+ test_ids = org_first_dates[org_id_col].tail(test_org_count).to_numpy()
89
+ is_test = df[org_id_col].isin(test_ids)
90
+ if is_test.all() or not is_test.any():
85
91
  continue
86
-
87
- # Calculate default rates
88
- train_rate = temp_train[label_col].mean()
89
- test_rate = temp_test[label_col].mean()
90
- rate_diff = abs(train_rate - test_rate)
91
-
92
- # Check if this split is better
92
+ rate_diff = abs(
93
+ df.loc[~is_test, label_col].mean() - df.loc[is_test, label_col].mean()
94
+ )
93
95
  if rate_diff < best_rate_diff:
94
- best_rate_diff = rate_diff
95
- best_split = {
96
- "test_orgs_data": test_orgs_data,
97
- "train_rate": train_rate,
98
- "test_rate": test_rate,
99
- "rate_diff": rate_diff,
100
- "test_org_count": test_org_count,
101
- "actual_test_size": len(temp_test) / len(df),
102
- }
103
-
104
- # Stop if we found an acceptable split
96
+ best_count, best_rate_diff = test_org_count, rate_diff
105
97
  if rate_diff <= rate_tolerance:
106
98
  print(
107
- f"✓ Found balanced split at {test_org_count} test orgs (rate diff: {rate_diff:.4f})"
99
+ f"Found a balanced split at {test_org_count} test entities "
100
+ f"(event-rate difference: {rate_diff:.4f})"
108
101
  )
109
102
  break
103
+ else:
104
+ if best_count is None:
105
+ best_count = n_test_orgs
106
+ print("Could not find any valid split; using the target size.")
107
+ else:
108
+ print(
109
+ f"Best achievable event-rate difference: {best_rate_diff:.4f} "
110
+ f"(exceeds the tolerance {rate_tolerance:.4f}); using "
111
+ f"{best_count} test entities."
112
+ )
110
113
 
111
- if best_split is None:
112
- # Fallback to original target
113
- best_split = {
114
- "test_orgs_data": org_first_dates.tail(n_test_orgs),
115
- "test_org_count": n_test_orgs,
116
- }
117
- print("⚠️ Could not find any valid split. Using target split.")
118
- elif best_rate_diff > rate_tolerance:
119
- print(
120
- f"⚠️ Best achievable rate difference: {best_rate_diff:.4f} (exceeds tolerance: {rate_tolerance:.4f})"
121
- )
122
-
123
- # Use the best split found
124
- test_orgs_data = best_split["test_orgs_data"]
125
-
126
- # Split data into train and test
127
- train_df = df[~df[org_id_col].isin(test_orgs_data[org_id_col].values)].copy()
128
- test_df = df[df[org_id_col].isin(test_orgs_data[org_id_col].values)].copy()
129
-
130
- # Calculate statistics
131
- total_orgs = df[org_id_col].nunique()
132
- total_records = len(df)
133
- total_pos = df[df[label_col] == 1].shape[0]
134
- total_neg = df[df[label_col] == 0].shape[0]
135
-
136
- train_orgs = train_df[org_id_col].nunique()
137
- train_records = len(train_df)
138
- train_pos = train_df[train_df[label_col] == 1].shape[0]
139
- train_neg = train_df[train_df[label_col] == 0].shape[0]
140
-
141
- test_orgs = test_df[org_id_col].nunique()
142
- test_records = len(test_df)
143
- test_pos = test_df[test_df[label_col] == 1].shape[0]
144
- test_neg = test_df[test_df[label_col] == 0].shape[0]
145
-
146
- # Verify no overlap
147
- train_org_set = set(train_df[org_id_col].unique())
148
- test_org_set = set(test_df[org_id_col].unique())
149
- overlap = train_org_set.intersection(test_org_set)
150
-
151
- # Print statistics
152
- print("\n=== FINAL SPLIT STATISTICS ===")
153
- print(f"Test organizations: {test_orgs} ({test_orgs / total_orgs:.1%})")
154
- print(f"Train organizations: {train_orgs} ({train_orgs / total_orgs:.1%})")
155
-
156
- print(f"\nTotal records: {total_records}")
157
- print(f"Train records: {train_records} ({train_records / total_records:.2%})")
158
- print(f"Test records: {test_records} ({test_records / total_records:.2%})")
159
-
160
- print("\nPositive samples:")
161
- print(f"Total: {total_pos}")
162
- print(f"Train: {train_pos} ({train_pos / total_pos:.2%})")
163
- print(f"Test: {test_pos} ({test_pos / total_pos:.2%})")
164
-
165
- print("\nNegative samples:")
166
- print(f"Total: {total_neg}")
167
- print(f"Train: {train_neg} ({train_neg / total_neg:.2%})")
168
- print(f"Test: {test_neg} ({test_neg / total_neg:.2%})")
169
-
170
- print("\nDefault rates:")
171
- print(f"Overall: {df[label_col].mean():.4f}")
172
- print(f"Train: {train_df[label_col].mean():.4f}")
173
- print(f"Test: {test_df[label_col].mean():.4f}")
174
- print(
175
- f"Difference: {abs(train_df[label_col].mean() - test_df[label_col].mean()):.4f}"
176
- )
177
-
178
- print(f"\nOrganization overlap: {len(overlap)} organizations (should be 0)")
114
+ test_orgs_data = org_first_dates.tail(best_count)
115
+ is_test = df[org_id_col].isin(test_orgs_data[org_id_col].to_numpy())
116
+ train_df = df[~is_test].copy()
117
+ test_df = df[is_test].copy()
179
118
 
180
- print("\nDate ranges:")
181
- print(
182
- "Cut-off date for splitting (earliest test org first loan):",
183
- test_orgs_data[date_col].min().strftime("%Y-%m-%d"),
184
- )
185
- print(
186
- f"Training orgs (first loan): {train_df.groupby(org_id_col)[date_col].min().min().strftime('%Y-%m-%d')} to {train_df.groupby(org_id_col)[date_col].min().max().strftime('%Y-%m-%d')}"
187
- )
188
- print(
189
- f"Testing orgs (first loan): {test_df.groupby(org_id_col)[date_col].min().min().strftime('%Y-%m-%d')} to {test_df.groupby(org_id_col)[date_col].min().max().strftime('%Y-%m-%d')}"
190
- )
191
- print(
192
- f"Training all observations: {train_df[date_col].min().strftime('%Y-%m-%d')} to {train_df[date_col].max().strftime('%Y-%m-%d')}"
119
+ stats = _calculate_split_statistics(
120
+ df, train_df, test_df, label_col, org_id_col, date_col
193
121
  )
122
+ overlap = set(train_df[org_id_col]) & set(test_df[org_id_col])
123
+ print(f"\nEntity overlap: {len(overlap)} (should be 0)")
194
124
  print(
195
- f"Testing all observations: {test_df[date_col].min().strftime('%Y-%m-%d')} to {test_df[date_col].max().strftime('%Y-%m-%d')}"
196
- )
197
-
198
- return (
199
- train_df,
200
- test_df,
201
- (
202
- train_df.shape,
203
- test_df.shape,
204
- train_df[train_df[label_col] == 1].shape,
205
- test_df[test_df[label_col] == 1].shape,
206
- ),
125
+ "Cut-off (first observation of the earliest test entity): "
126
+ f"{test_orgs_data[date_col].min():%Y-%m-%d}"
207
127
  )
128
+ return train_df, test_df, stats
208
129
 
209
130
 
210
131
  def split_data_stratified_by_org(
@@ -2,7 +2,7 @@
2
2
 
3
3
  import math
4
4
  from collections.abc import Sequence
5
- from typing import Union
5
+ from typing import Optional, Union
6
6
 
7
7
  import matplotlib.pyplot as plt
8
8
  import numpy as np
@@ -45,28 +45,47 @@ def analyze_shap_interactions(
45
45
  shap_values: np.ndarray,
46
46
  n_important_features: int = 2,
47
47
  class_index: int = -1,
48
+ features: Optional[Sequence[str]] = None,
48
49
  ) -> plt.Figure:
49
50
  """Grid of main effects (diagonal) and pairwise interactions (off-diagonal).
50
51
 
52
+ The grid shows the `n_important_features` features with the highest mean
53
+ absolute SHAP value (from `shap_values`), most important first, or the
54
+ explicit `features` when given. It used to show the first features of
55
+ `feature_names`, whatever their importance.
56
+
51
57
  ``class_index`` selects which class's interaction values to plot when the
52
58
  model's SHAP output carries a per-class axis; ``shap_values`` must be the
53
- matching 2D slice for the same class.
59
+ matching 2D slice for the same class, with one column per
60
+ `feature_names` entry.
54
61
  """
62
+ feature_list = list(feature_names)
63
+ if features is not None:
64
+ unknown = [f for f in features if f not in feature_list]
65
+ if unknown:
66
+ raise ValueError(f"features not in feature_names: {unknown}")
67
+ important_features = list(features)
68
+ else:
69
+ values = np.asarray(shap_values)
70
+ if values.ndim != 2 or values.shape[1] != len(feature_list):
71
+ raise ValueError(
72
+ "shap_values must be 2D with one column per feature_names entry, "
73
+ f"got shape {values.shape} for {len(feature_list)} features."
74
+ )
75
+ ranking = np.argsort(-np.abs(values).mean(axis=0), kind="stable")
76
+ important_features = [feature_list[i] for i in ranking[:n_important_features]]
77
+ n_features = len(important_features)
78
+
55
79
  interaction_values = compute_shap_interaction_values(
56
80
  model, X_test, feature_names, class_index=class_index
57
81
  )
58
82
 
59
- important_features = list(feature_names)[:n_important_features]
60
- n_features = len(important_features)
61
-
62
- fig, axes = plt.subplots(
63
- n_features, n_features, figsize=(10 * n_features, 10 * n_features)
64
- )
83
+ side = 5 * n_features # 5 inches per panel
84
+ fig, axes = plt.subplots(n_features, n_features, figsize=(side, side))
65
85
 
66
86
  if n_features == 1:
67
87
  axes = np.array([[axes]])
68
88
 
69
- feature_list = list(feature_names)
70
89
  for i in range(n_features):
71
90
  for j in range(n_features):
72
91
  feature_i = important_features[i]
@@ -111,6 +130,10 @@ def plot_top_shap_interactions(
111
130
  ) -> plt.Figure:
112
131
  """Plot the strongest feature interactions by mean absolute SHAP interaction.
113
132
 
133
+ Each unordered pair appears once: the interaction matrix is symmetric, so
134
+ only its upper triangle is ranked (it used to rank both triangles, so
135
+ every pair was plotted twice).
136
+
114
137
  ``class_index`` selects which class's interaction values to rank and plot
115
138
  when the model's SHAP output carries a per-class axis.
116
139
  """
@@ -122,18 +145,15 @@ def plot_top_shap_interactions(
122
145
  )
123
146
 
124
147
  interaction_strength = np.abs(interaction_values).mean(axis=0)
125
- np.fill_diagonal(interaction_strength, 0)
126
-
127
- flat_indices = np.argsort(interaction_strength.flatten())[-n_top_interactions:]
128
- feature_pairs: list[tuple] = []
129
148
  names = list(feature_names)
130
149
 
131
- for flat_idx in flat_indices[::-1]:
132
- i, j = np.unravel_index(flat_idx, interaction_strength.shape)
133
- if i < j:
134
- feature_pairs.append((names[i], names[j], interaction_strength[i, j]))
135
- else:
136
- feature_pairs.append((names[j], names[i], interaction_strength[i, j]))
150
+ # Symmetric matrix: rank each unordered pair once (upper triangle).
151
+ rows, cols = np.triu_indices(len(names), k=1)
152
+ strengths = interaction_strength[rows, cols]
153
+ order = np.argsort(-strengths, kind="stable")[:n_top_interactions]
154
+ feature_pairs: list[tuple] = [
155
+ (names[rows[k]], names[cols[k]], float(strengths[k])) for k in order
156
+ ]
137
157
 
138
158
  n_cols = 2
139
159
  n_rows = math.ceil(len(feature_pairs) / n_cols)
@@ -60,6 +60,15 @@ class ExplainabilityToolkit:
60
60
  self.feature_names = list(feature_names)
61
61
  self.numerical_features = list(numerical_features or [])
62
62
  self.categorical_features = list(categorical_features or [])
63
+ unknown = [
64
+ f
65
+ for f in self.numerical_features + self.categorical_features
66
+ if f not in self.feature_names
67
+ ]
68
+ if unknown:
69
+ raise ValueError(
70
+ f"numerical/categorical features not in feature_names: {unknown}"
71
+ )
63
72
  self.explainer = TreeSHAPExplainer(
64
73
  model,
65
74
  cat_features=cat_features or categorical_features,
@@ -128,13 +137,14 @@ class ExplainabilityToolkit:
128
137
  if not self.numerical_features:
129
138
  return None
130
139
  self._require_pos_label_if_multiclass("plot_numerical_dependence")
131
- n_num = len(self.numerical_features)
132
140
  values_2d = (
133
141
  self.shap_result_.values_for_label(self.pos_label)
134
142
  if self.shap_result_.is_multiclass
135
143
  else self.shap_result_.values_for_class()
136
144
  )
137
- shap_subset = values_2d[:, :n_num]
145
+ # Select the SHAP columns by name: feature_names need not list the
146
+ # numerical features first.
147
+ shap_subset = values_2d[:, self._columns(self.numerical_features)]
138
148
  X_num = self.shap_result_.data[self.numerical_features]
139
149
  return plot_shap_dependence_grid(
140
150
  feature_list=self.numerical_features,
@@ -156,13 +166,12 @@ class ExplainabilityToolkit:
156
166
  if not self.categorical_features:
157
167
  return None
158
168
  self._require_pos_label_if_multiclass("plot_categorical_dependence")
159
- n_num = len(self.numerical_features)
160
169
  values_2d = (
161
170
  self.shap_result_.values_for_label(self.pos_label)
162
171
  if self.shap_result_.is_multiclass
163
172
  else self.shap_result_.values_for_class()
164
173
  )
165
- shap_cat = values_2d[:, n_num:]
174
+ shap_cat = values_2d[:, self._columns(self.categorical_features)]
166
175
  X_cat = self.shap_result_.data[self.categorical_features]
167
176
  return plot_shap_dependence_grid(
168
177
  feature_list=self.categorical_features,
@@ -181,6 +190,10 @@ class ExplainabilityToolkit:
181
190
  return self.shap_result_.class_index_for_label(self.pos_label)
182
191
  return -1
183
192
 
193
+ def _columns(self, features: Sequence[str]) -> list[int]:
194
+ """Positions of `features` in the SHAP matrix (``feature_names`` order)."""
195
+ return [self.feature_names.index(f) for f in features]
196
+
184
197
  def plot_interaction_grid(self, n_important_features: int = 2):
185
198
  """Main-effect / interaction grid for top features.
186
199