classifier-toolkit 0.3.5__tar.gz → 0.4.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.
Potentially problematic release.
This version of classifier-toolkit might be problematic. Click here for more details.
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/.gitignore +5 -1
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/PKG-INFO +49 -14
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/README.md +47 -13
- classifier_toolkit-0.4.0/classifier_toolkit/calibration/__init__.py +34 -0
- classifier_toolkit-0.4.0/classifier_toolkit/calibration/base.py +5 -0
- classifier_toolkit-0.4.0/classifier_toolkit/calibration/ovr_calibration.py +254 -0
- classifier_toolkit-0.4.0/classifier_toolkit/calibration/reliability.py +109 -0
- classifier_toolkit-0.4.0/classifier_toolkit/datasets/__init__.py +3 -0
- classifier_toolkit-0.4.0/classifier_toolkit/datasets/_demo.py +178 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/eda_toolkit.py +6 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/visualizations.py +314 -91
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/explainability/__init__.py +2 -1
- classifier_toolkit-0.4.0/classifier_toolkit/explainability/base.py +138 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/explainability/interactions.py +34 -6
- classifier_toolkit-0.4.0/classifier_toolkit/explainability/misclassification.py +244 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/explainability/plots.py +56 -7
- classifier_toolkit-0.4.0/classifier_toolkit/explainability/toolkit.py +234 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/explainability/tree_explainer.py +49 -1
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/counter_intuitive.py +65 -7
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/drift.py +29 -18
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/reducer.py +17 -5
- classifier_toolkit-0.4.0/classifier_toolkit/feature_selection/utils/scoring.py +375 -0
- classifier_toolkit-0.4.0/classifier_toolkit/feature_selection/wrapper_methods/combination_search.py +1131 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/recursive_feature_eliminator.py +538 -18
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/rfe.py +64 -29
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/hyper_parameter_tuning/tuner.py +356 -60
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/models/ensemble_methods.py +93 -30
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/utils/params.py +45 -4
- classifier_toolkit-0.4.0/classifier_toolkit/model_validation/__init__.py +27 -0
- classifier_toolkit-0.4.0/classifier_toolkit/model_validation/evaluator.py +1089 -0
- classifier_toolkit-0.4.0/classifier_toolkit/risk_class/__init__.py +37 -0
- classifier_toolkit-0.4.0/classifier_toolkit/risk_class/risk_classes.py +1144 -0
- classifier_toolkit-0.4.0/classifier_toolkit/risk_class/risk_classes_dp.py +670 -0
- classifier_toolkit-0.4.0/examples/example_bayesian_search.ipynb +492 -0
- classifier_toolkit-0.4.0/examples/example_combination_feature_search.ipynb +404 -0
- classifier_toolkit-0.4.0/examples/example_explainability_catboost.ipynb +641 -0
- classifier_toolkit-0.4.0/examples/example_explainability_lgbm.ipynb +672 -0
- classifier_toolkit-0.4.0/examples/example_feature_reduction.ipynb +376 -0
- classifier_toolkit-0.4.0/examples/example_grid_search.ipynb +588 -0
- classifier_toolkit-0.4.0/examples/example_model_training_catboost.ipynb +194 -0
- classifier_toolkit-0.4.0/examples/example_model_training_lgbm.ipynb +186 -0
- classifier_toolkit-0.4.0/examples/example_model_validation.ipynb +369 -0
- classifier_toolkit-0.4.0/examples/example_recursive_feature_eliminator.ipynb +539 -0
- classifier_toolkit-0.4.0/examples/example_risk_classes_dp.ipynb +290 -0
- classifier_toolkit-0.4.0/examples/example_train_test_partition.ipynb +276 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/pyproject.toml +16 -1
- classifier_toolkit-0.4.0/tests/calibration/test_ovr_calibration.py +271 -0
- classifier_toolkit-0.4.0/tests/calibration/test_reliability.py +67 -0
- classifier_toolkit-0.4.0/tests/datasets/test_demo_data.py +65 -0
- classifier_toolkit-0.4.0/tests/eda/test_visualizations.py +400 -0
- classifier_toolkit-0.4.0/tests/explainability/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/explainability/test_catboost_e2e.py +2 -3
- classifier_toolkit-0.4.0/tests/explainability/test_interactions.py +40 -0
- classifier_toolkit-0.4.0/tests/explainability/test_multiclass_e2e.py +473 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_counter_intuitive.py +104 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_drift.py +112 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_reducer.py +50 -0
- classifier_toolkit-0.4.0/tests/feature_selection/test_combination_search.py +842 -0
- classifier_toolkit-0.4.0/tests/feature_selection/test_recursive_feature_eliminator.py +1537 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_rfe.py +66 -6
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_scoring.py +137 -6
- classifier_toolkit-0.4.0/tests/model_training/__init__.py +0 -0
- classifier_toolkit-0.4.0/tests/model_training/hyper_parameter_tuning/__init__.py +0 -0
- classifier_toolkit-0.4.0/tests/model_training/hyper_parameter_tuning/test_params.py +103 -0
- classifier_toolkit-0.4.0/tests/model_training/hyper_parameter_tuning/test_tuner.py +877 -0
- classifier_toolkit-0.4.0/tests/model_training/models/__init__.py +0 -0
- classifier_toolkit-0.4.0/tests/model_training/models/test_ensemble_methods.py +160 -0
- classifier_toolkit-0.4.0/tests/model_validation/__init__.py +0 -0
- classifier_toolkit-0.4.0/tests/model_validation/test_compare_score_distributions.py +90 -0
- classifier_toolkit-0.4.0/tests/model_validation/test_evaluate_risk_classes_dp.py +46 -0
- classifier_toolkit-0.4.0/tests/model_validation/test_evaluator_threshold.py +33 -0
- classifier_toolkit-0.4.0/tests/model_validation/test_print_full_validation_report.py +101 -0
- classifier_toolkit-0.4.0/tests/model_validation/test_print_risk_class_validation_report.py +71 -0
- classifier_toolkit-0.4.0/tests/model_validation/test_score_distribution_psi.py +69 -0
- classifier_toolkit-0.4.0/tests/risk_class/__init__.py +0 -0
- classifier_toolkit-0.4.0/tests/risk_class/test_construct_bins_dp.py +281 -0
- classifier_toolkit-0.4.0/tests/risk_class/test_risk_classes.py +193 -0
- classifier_toolkit-0.4.0/tests/risk_class/test_validate_risk_classes.py +137 -0
- classifier_toolkit-0.3.5/.github/pull_request_template/default.md +0 -13
- classifier_toolkit-0.3.5/.github/workflows/checks.yaml +0 -134
- classifier_toolkit-0.3.5/.github/workflows/docs.yml +0 -31
- classifier_toolkit-0.3.5/.github/workflows/master.yaml +0 -48
- classifier_toolkit-0.3.5/.github/workflows/release.yaml +0 -51
- classifier_toolkit-0.3.5/.github/workflows/working-branch.yaml +0 -14
- classifier_toolkit-0.3.5/.python-version +0 -1
- classifier_toolkit-0.3.5/.sqlfluff +0 -38
- classifier_toolkit-0.3.5/Makefile +0 -22
- classifier_toolkit-0.3.5/classifier_toolkit/explainability/base.py +0 -60
- classifier_toolkit-0.3.5/classifier_toolkit/explainability/misclassification.py +0 -107
- classifier_toolkit-0.3.5/classifier_toolkit/explainability/toolkit.py +0 -137
- classifier_toolkit-0.3.5/classifier_toolkit/feature_selection/utils/scoring.py +0 -226
- classifier_toolkit-0.3.5/classifier_toolkit/feature_selection/wrapper_methods/combination_search.py +0 -569
- classifier_toolkit-0.3.5/docs/CNAME +0 -1
- classifier_toolkit-0.3.5/docs/changelog.md +0 -36
- classifier_toolkit-0.3.5/docs/data_partition/data_preprocess.md +0 -43
- classifier_toolkit-0.3.5/docs/data_partition/optimize_data.md +0 -30
- classifier_toolkit-0.3.5/docs/data_partition/overview.md +0 -30
- classifier_toolkit-0.3.5/docs/data_partition/split_train_test.md +0 -43
- classifier_toolkit-0.3.5/docs/eda/bivariate_analysis.md +0 -38
- classifier_toolkit-0.3.5/docs/eda/eda_toolkit.md +0 -65
- classifier_toolkit-0.3.5/docs/eda/feature_engineering.md +0 -60
- classifier_toolkit-0.3.5/docs/eda/first_glance.md +0 -54
- classifier_toolkit-0.3.5/docs/eda/overview.md +0 -128
- classifier_toolkit-0.3.5/docs/eda/univariate_analysis.md +0 -48
- classifier_toolkit-0.3.5/docs/eda/visualizations.md +0 -51
- classifier_toolkit-0.3.5/docs/eda/warnings/default_warnings.md +0 -45
- classifier_toolkit-0.3.5/docs/eda/warnings/warning_system.md +0 -32
- classifier_toolkit-0.3.5/docs/examples/eda_example.md +0 -71
- classifier_toolkit-0.3.5/docs/examples/feature_selection_advanced.md +0 -111
- classifier_toolkit-0.3.5/docs/examples/feature_selection_example.md +0 -123
- classifier_toolkit-0.3.5/docs/explainability/interactions.md +0 -48
- classifier_toolkit-0.3.5/docs/explainability/misclassification.md +0 -46
- classifier_toolkit-0.3.5/docs/explainability/overview.md +0 -87
- classifier_toolkit-0.3.5/docs/explainability/plots.md +0 -47
- classifier_toolkit-0.3.5/docs/explainability/toolkit.md +0 -59
- classifier_toolkit-0.3.5/docs/explainability/tree_explainer.md +0 -50
- classifier_toolkit-0.3.5/docs/feature_reduction/correlation.md +0 -34
- classifier_toolkit-0.3.5/docs/feature_reduction/counter_intuitive.md +0 -49
- classifier_toolkit-0.3.5/docs/feature_reduction/drift.md +0 -35
- classifier_toolkit-0.3.5/docs/feature_reduction/expert_rules.md +0 -26
- classifier_toolkit-0.3.5/docs/feature_reduction/low_variance.md +0 -22
- classifier_toolkit-0.3.5/docs/feature_reduction/overview.md +0 -49
- classifier_toolkit-0.3.5/docs/feature_reduction/predictive_power.md +0 -55
- classifier_toolkit-0.3.5/docs/feature_reduction/reducer.md +0 -46
- classifier_toolkit-0.3.5/docs/feature_selection/embedded_methods/elastic_net.md +0 -54
- classifier_toolkit-0.3.5/docs/feature_selection/feature_stability.md +0 -38
- classifier_toolkit-0.3.5/docs/feature_selection/meta_selector.md +0 -99
- classifier_toolkit-0.3.5/docs/feature_selection/overview.md +0 -89
- classifier_toolkit-0.3.5/docs/feature_selection/utils/data_handling.md +0 -60
- classifier_toolkit-0.3.5/docs/feature_selection/utils/scoring.md +0 -50
- classifier_toolkit-0.3.5/docs/feature_selection/wrapper_methods/bayesian_search.md +0 -44
- classifier_toolkit-0.3.5/docs/feature_selection/wrapper_methods/boruta.md +0 -49
- classifier_toolkit-0.3.5/docs/feature_selection/wrapper_methods/combination_search.md +0 -39
- classifier_toolkit-0.3.5/docs/feature_selection/wrapper_methods/recursive_feature_eliminator.md +0 -35
- classifier_toolkit-0.3.5/docs/feature_selection/wrapper_methods/rfe.md +0 -90
- classifier_toolkit-0.3.5/docs/feature_selection/wrapper_methods/sequential_selection.md +0 -54
- classifier_toolkit-0.3.5/docs/index.md +0 -92
- classifier_toolkit-0.3.5/docs/model_training/overview.md +0 -51
- classifier_toolkit-0.3.5/docs/model_training/tuner.md +0 -359
- classifier_toolkit-0.3.5/docs/reference/data_partition/data_preprocess.md +0 -46
- classifier_toolkit-0.3.5/docs/reference/data_partition/optimize_data.md +0 -36
- classifier_toolkit-0.3.5/docs/reference/data_partition/overview.md +0 -37
- classifier_toolkit-0.3.5/docs/reference/data_partition/split_train_test.md +0 -57
- classifier_toolkit-0.3.5/docs/reference/eda/bivariate_analysis.md +0 -42
- classifier_toolkit-0.3.5/docs/reference/eda/eda_toolkit.md +0 -47
- classifier_toolkit-0.3.5/docs/reference/eda/feature_engineering.md +0 -42
- classifier_toolkit-0.3.5/docs/reference/eda/first_glance.md +0 -42
- classifier_toolkit-0.3.5/docs/reference/eda/overview.md +0 -35
- classifier_toolkit-0.3.5/docs/reference/eda/univariate_analysis.md +0 -43
- classifier_toolkit-0.3.5/docs/reference/eda/visualizations.md +0 -39
- classifier_toolkit-0.3.5/docs/reference/eda/warnings/default_warnings.md +0 -52
- classifier_toolkit-0.3.5/docs/reference/eda/warnings/warning_system.md +0 -36
- classifier_toolkit-0.3.5/docs/reference/explainability/interactions.md +0 -7
- classifier_toolkit-0.3.5/docs/reference/explainability/misclassification.md +0 -3
- classifier_toolkit-0.3.5/docs/reference/explainability/overview.md +0 -42
- classifier_toolkit-0.3.5/docs/reference/explainability/plots.md +0 -7
- classifier_toolkit-0.3.5/docs/reference/explainability/toolkit.md +0 -3
- classifier_toolkit-0.3.5/docs/reference/explainability/tree_explainer.md +0 -7
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/base.md +0 -5
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/correlation.md +0 -49
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/counter_intuitive.md +0 -40
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/drift.md +0 -75
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/expert_rules.md +0 -16
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/low_variance.md +0 -25
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/overview.md +0 -33
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/predictive_power.md +0 -68
- classifier_toolkit-0.3.5/docs/reference/feature_reduction/reducer.md +0 -121
- classifier_toolkit-0.3.5/docs/reference/feature_selection/base.md +0 -3
- classifier_toolkit-0.3.5/docs/reference/feature_selection/embedded_methods/elastic_net.md +0 -40
- classifier_toolkit-0.3.5/docs/reference/feature_selection/feature_stability.md +0 -38
- classifier_toolkit-0.3.5/docs/reference/feature_selection/meta_selector.md +0 -41
- classifier_toolkit-0.3.5/docs/reference/feature_selection/overview.md +0 -28
- classifier_toolkit-0.3.5/docs/reference/feature_selection/utils/data_handling.md +0 -43
- classifier_toolkit-0.3.5/docs/reference/feature_selection/utils/plottings.md +0 -5
- classifier_toolkit-0.3.5/docs/reference/feature_selection/utils/scoring.md +0 -37
- classifier_toolkit-0.3.5/docs/reference/feature_selection/wrapper_methods/bayesian_search.md +0 -35
- classifier_toolkit-0.3.5/docs/reference/feature_selection/wrapper_methods/boruta.md +0 -34
- classifier_toolkit-0.3.5/docs/reference/feature_selection/wrapper_methods/combination_search.md +0 -111
- classifier_toolkit-0.3.5/docs/reference/feature_selection/wrapper_methods/recursive_feature_eliminator.md +0 -268
- classifier_toolkit-0.3.5/docs/reference/feature_selection/wrapper_methods/rfe.md +0 -62
- classifier_toolkit-0.3.5/docs/reference/feature_selection/wrapper_methods/sequential_selection.md +0 -55
- classifier_toolkit-0.3.5/docs/reference/model_training/params.md +0 -5
- classifier_toolkit-0.3.5/docs/reference/tuner/tuner.md +0 -3
- classifier_toolkit-0.3.5/examples/__init__.py +0 -1
- classifier_toolkit-0.3.5/examples/example_bayesian_search.ipynb +0 -4475
- classifier_toolkit-0.3.5/examples/example_combination_feature_search.ipynb +0 -4262
- classifier_toolkit-0.3.5/examples/example_explainability_catboost.ipynb +0 -1565
- classifier_toolkit-0.3.5/examples/example_explainability_lgbm.ipynb +0 -1625
- classifier_toolkit-0.3.5/examples/example_feature_reduction.ipynb +0 -2280
- classifier_toolkit-0.3.5/examples/example_grid_search.ipynb +0 -1975
- classifier_toolkit-0.3.5/examples/example_model_training_catboost.ipynb +0 -133
- classifier_toolkit-0.3.5/examples/example_model_training_lgbm.ipynb +0 -133
- classifier_toolkit-0.3.5/examples/example_recursive_feature_eliminator.ipynb +0 -862
- classifier_toolkit-0.3.5/examples/example_train_test_partition.ipynb +0 -446
- classifier_toolkit-0.3.5/main.py +0 -6
- classifier_toolkit-0.3.5/mkdocs.yml +0 -195
- classifier_toolkit-0.3.5/notebooks/paylater_removed.json +0 -244
- classifier_toolkit-0.3.5/ruff.toml +0 -46
- classifier_toolkit-0.3.5/tests/eda/test_visualizations.py +0 -116
- classifier_toolkit-0.3.5/tests/feature_selection/test_combination_search.py +0 -290
- classifier_toolkit-0.3.5/tests/feature_selection/test_recursive_feature_eliminator.py +0 -638
- classifier_toolkit-0.3.5/uv.lock +0 -3998
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/LICENSE +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/data_partition/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/data_partition/data_preprocess.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/data_partition/optimize_data.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/data_partition/split_train_test.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/bivariate_analysis.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/feature_engineering.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/first_glance.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/univariate_analysis.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/warnings/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/warnings/automated_warnings.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/eda/warnings/default_warnings.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/base.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/correlation.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/expert_rules.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/low_variance.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_reduction/predictive_power.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/base.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/embedded_methods/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/embedded_methods/elastic_net.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/feature_stability.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/meta_selector.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/utils/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/utils/data_handling.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/utils/feature_type_constraints.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/utils/plottings.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/bayesian_search.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/boruta.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/rfe_catboost.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/feature_selection/wrapper_methods/sequential_selection.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/hyper_parameter_tuning/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/models/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/models/base.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/classifier_toolkit/model_training/utils/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/__init__.py +0 -0
- {classifier_toolkit-0.3.5/tests/data_partition → classifier_toolkit-0.4.0/tests/calibration}/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/conftest.py +0 -0
- {classifier_toolkit-0.3.5/tests/eda → classifier_toolkit-0.4.0/tests/data_partition}/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/data_partition/test_data_preprocess.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/data_partition/test_optimize_data.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/data_partition/test_split_train_test.py +0 -0
- {classifier_toolkit-0.3.5/tests/explainability → classifier_toolkit-0.4.0/tests/datasets}/__init__.py +0 -0
- /classifier_toolkit-0.3.5/docs/stylesheets/extra.css → /classifier_toolkit-0.4.0/tests/eda/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/eda/test_bivariate_analysis.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/eda/test_feature_engineering.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/eda/test_first_glance.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/eda/test_univariate_analysis.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/eda/test_warnings.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/explainability/test_misclassification.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/explainability/test_plots.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/explainability/test_smoke.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/explainability/test_toolkit.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/explainability/test_tree_explainer.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_correlation.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_expert_rules.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_low_variance.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_predictive_power.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_reduction/test_smoke.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/__init__.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_bayesian_search.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_boruta.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_elastic_net.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_feature_stability.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_rfe_catboost.py +0 -0
- {classifier_toolkit-0.3.5 → classifier_toolkit-0.4.0}/tests/feature_selection/test_sequential_selection.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: classifier-toolkit
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.4.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
|
|
@@ -18,6 +18,7 @@ Requires-Dist: polars<2.0.0,>=1.2.1
|
|
|
18
18
|
Requires-Dist: pyarrow>=18.0.0
|
|
19
19
|
Requires-Dist: scikit-learn<2.0.0,>=1.4.0
|
|
20
20
|
Requires-Dist: scipy>=1.11.0
|
|
21
|
+
Requires-Dist: seaborn<0.14.0,>=0.13.0
|
|
21
22
|
Requires-Dist: shap>=0.46.0
|
|
22
23
|
Requires-Dist: statsmodels<0.15.0,>=0.14.2
|
|
23
24
|
Requires-Dist: tabulate<0.10.0,>=0.9.0
|
|
@@ -60,13 +61,17 @@ This library is published in the PyPI directory. To install, users can run pip i
|
|
|
60
61
|
|
|
61
62
|
### Usage
|
|
62
63
|
|
|
63
|
-
This library automates binary classification
|
|
64
|
+
This library automates binary and multiclass classification workflows. It is independent of the modelled problem: the class of interest is configured through `pos_label` (the positive class for binary targets, the class of interest for multiclass ones). It includes several packages designed to address the main steps in any machine learning/data science task:
|
|
64
65
|
|
|
65
|
-
1. **EDA**: accessible via `
|
|
66
|
+
1. **EDA**: accessible via `EDAToolkit`. Provides EDA and feature engineering functionality with all necessary visualizations.
|
|
66
67
|
2. **Feature Reduction**: filter-style pre-selection pipeline (expert rules, low variance, drift, predictive power, counter-intuitive direction, high correlation).
|
|
67
68
|
3. **Feature Selection**: wrapper and embedded methods (RFE, Boruta, Sequential, Bayesian, ElasticNet, MetaSelector).
|
|
68
|
-
4. **Model Training**: accessible via `Tuner`. Hyperparameter optimization (grid search, Bayesian via Optuna) for LightGBM and CatBoost with train/val/test evaluation.
|
|
69
|
-
5.
|
|
69
|
+
4. **Model Training**: accessible via `Tuner`. Hyperparameter optimization (grid search, Bayesian via Optuna) for LightGBM and CatBoost with train/val/test evaluation, including multiclass objectives and class weights.
|
|
70
|
+
5. **Risk Class**: `RiskClassBuilderDP` / `RiskClassBuilder` turn model scores into risk classes with statistically validated, ordered event rates.
|
|
71
|
+
6. **Model Validation**: `Evaluator` computes metrics, calibration and stability checks for any model exposing `predict_proba`.
|
|
72
|
+
7. **Calibration**: `OVRHistogramCalibrator` and `plot_reliability_curves` for One-vs-Rest probability calibration.
|
|
73
|
+
8. **Explainability**: SHAP-based explanations, interactions and misclassification diagnostics.
|
|
74
|
+
9. **Data Partition**: temporal-aware train/test splitting, preprocessing and dtype optimization.
|
|
70
75
|
|
|
71
76
|
For detailed usage, refer to the documentation.
|
|
72
77
|
|
|
@@ -139,16 +144,46 @@ For detailed usage, refer to the documentation.
|
|
|
139
144
|
- **Model Training**: Hyperparameter optimization for LightGBM and CatBoost, with support for grid search and Bayesian optimization (via Optuna).
|
|
140
145
|
|
|
141
146
|
```python
|
|
142
|
-
from classifier_toolkit.model_training.hyper_parameter_tuning import Tuner
|
|
147
|
+
from classifier_toolkit.model_training.hyper_parameter_tuning.tuner import Tuner
|
|
148
|
+
|
|
149
|
+
tuner = Tuner(
|
|
150
|
+
X=X_train, y=y_train,
|
|
151
|
+
model_name="lightgbm",
|
|
152
|
+
X_val=X_val, y_val=y_val,
|
|
153
|
+
X_test=X_test, y_test=y_test,
|
|
154
|
+
search_method="bayesian",
|
|
155
|
+
n_trials=50,
|
|
156
|
+
optimization_metric="prauc",
|
|
157
|
+
)
|
|
158
|
+
result = tuner.tune()
|
|
159
|
+
|
|
160
|
+
best_model = result["best_model"]
|
|
161
|
+
result["trials_results"] # full trial results (DataFrame)
|
|
162
|
+
```
|
|
163
|
+
|
|
164
|
+
Reported metrics are `auc`, `prauc`, `ks`, `log_loss` and `brier` (`ks`/`brier` are binary-only). Custom parameter search spaces can be defined via `ModelParams` and `ParamRange`.
|
|
165
|
+
|
|
166
|
+
- **Risk Class**: Builds risk classes from model scores. `RiskClassBuilderDP` searches bin edges with dynamic programming so that each class' observed event rate falls in a target band (`target_ranges`, required), then checks that adjacent classes are statistically distinguishable; `RiskClassBuilder` discovers classes with KMeans.
|
|
143
167
|
|
|
144
|
-
|
|
145
|
-
|
|
168
|
+
```python
|
|
169
|
+
from classifier_toolkit.risk_class import RiskClassBuilderDP
|
|
146
170
|
|
|
147
|
-
|
|
148
|
-
|
|
171
|
+
builder = RiskClassBuilderDP(
|
|
172
|
+
target_col="target",
|
|
173
|
+
target_ranges=[(0.00, 0.02), (0.02, 0.05), (0.05, 0.10)],
|
|
174
|
+
min_obs_per_bin=200,
|
|
175
|
+
)
|
|
176
|
+
result = builder.build(train_proba, y_train)
|
|
177
|
+
print(result.bins, result.n_classes)
|
|
149
178
|
```
|
|
150
179
|
|
|
151
|
-
|
|
180
|
+
- **Model Validation**: `Evaluator` evaluates any model with `predict_proba` (metrics, ROC/PR/calibration/threshold plots), `score_distribution_psi` and `compare_score_distributions` compare a reference score distribution with a current one, and `evaluate_risk_classes_dp` checks calibration within each risk class.
|
|
181
|
+
|
|
182
|
+
- **Calibration**: `OVRHistogramCalibrator` calibrates multiclass (or binary) probabilities One-vs-Rest with histogram binning; `plot_reliability_curves` shows raw vs calibrated reliability per class.
|
|
183
|
+
|
|
184
|
+
- **Explainability**: `TreeSHAPExplainer`, `ExplainabilityToolkit`, SHAP plots, pairwise interaction analysis and a `MisclassificationAnalyzer` that explains confusion-matrix quadrants (per-class SHAP for multiclass models).
|
|
185
|
+
|
|
186
|
+
- **Data Partition**: Temporal-aware train/test splitting, preprocessing helpers and dtype optimization.
|
|
152
187
|
|
|
153
188
|
### Development & CI/CD
|
|
154
189
|
|
|
@@ -162,8 +197,8 @@ This project uses modern tooling for fast and efficient development workflows:
|
|
|
162
197
|
#### CI/CD Pipeline
|
|
163
198
|
Our CI/CD pipeline is optimized for speed and efficiency:
|
|
164
199
|
|
|
165
|
-
- **Parallel Test Execution**:
|
|
166
|
-
- **Shared Caching**:
|
|
200
|
+
- **Parallel Test Execution**: One test job per directory under `tests/` (discovered automatically, so new test groups are picked up without editing the workflow), all running simultaneously
|
|
201
|
+
- **Shared Caching**: The parallel jobs share the same dependency cache, avoiding duplicate downloads
|
|
167
202
|
- **Smart Test Reruns**: Failed tests run first (`pytest --lf --ff`) for faster feedback on fixes
|
|
168
203
|
- **Master Protection**: Build tests only run on `master` branch and PRs targeting `master`, saving CI resources on feature branches
|
|
169
204
|
- **Automatic Linting**: Code quality checks (Ruff, SQLFluff) run on every push
|
|
@@ -188,5 +223,5 @@ uv build
|
|
|
188
223
|
|
|
189
224
|
### Future Work
|
|
190
225
|
The next planned improvements and additions to the library include:
|
|
191
|
-
*
|
|
226
|
+
* Extending the evaluation and reporting tools in `model_validation` (more checks and report formats).
|
|
192
227
|
* Expanding documentation to include architecture diagrams and detailed usage examples.
|
|
@@ -29,13 +29,17 @@ This library is published in the PyPI directory. To install, users can run pip i
|
|
|
29
29
|
|
|
30
30
|
### Usage
|
|
31
31
|
|
|
32
|
-
This library automates binary classification
|
|
32
|
+
This library automates binary and multiclass classification workflows. It is independent of the modelled problem: the class of interest is configured through `pos_label` (the positive class for binary targets, the class of interest for multiclass ones). It includes several packages designed to address the main steps in any machine learning/data science task:
|
|
33
33
|
|
|
34
|
-
1. **EDA**: accessible via `
|
|
34
|
+
1. **EDA**: accessible via `EDAToolkit`. Provides EDA and feature engineering functionality with all necessary visualizations.
|
|
35
35
|
2. **Feature Reduction**: filter-style pre-selection pipeline (expert rules, low variance, drift, predictive power, counter-intuitive direction, high correlation).
|
|
36
36
|
3. **Feature Selection**: wrapper and embedded methods (RFE, Boruta, Sequential, Bayesian, ElasticNet, MetaSelector).
|
|
37
|
-
4. **Model Training**: accessible via `Tuner`. Hyperparameter optimization (grid search, Bayesian via Optuna) for LightGBM and CatBoost with train/val/test evaluation.
|
|
38
|
-
5.
|
|
37
|
+
4. **Model Training**: accessible via `Tuner`. Hyperparameter optimization (grid search, Bayesian via Optuna) for LightGBM and CatBoost with train/val/test evaluation, including multiclass objectives and class weights.
|
|
38
|
+
5. **Risk Class**: `RiskClassBuilderDP` / `RiskClassBuilder` turn model scores into risk classes with statistically validated, ordered event rates.
|
|
39
|
+
6. **Model Validation**: `Evaluator` computes metrics, calibration and stability checks for any model exposing `predict_proba`.
|
|
40
|
+
7. **Calibration**: `OVRHistogramCalibrator` and `plot_reliability_curves` for One-vs-Rest probability calibration.
|
|
41
|
+
8. **Explainability**: SHAP-based explanations, interactions and misclassification diagnostics.
|
|
42
|
+
9. **Data Partition**: temporal-aware train/test splitting, preprocessing and dtype optimization.
|
|
39
43
|
|
|
40
44
|
For detailed usage, refer to the documentation.
|
|
41
45
|
|
|
@@ -108,16 +112,46 @@ For detailed usage, refer to the documentation.
|
|
|
108
112
|
- **Model Training**: Hyperparameter optimization for LightGBM and CatBoost, with support for grid search and Bayesian optimization (via Optuna).
|
|
109
113
|
|
|
110
114
|
```python
|
|
111
|
-
from classifier_toolkit.model_training.hyper_parameter_tuning import Tuner
|
|
115
|
+
from classifier_toolkit.model_training.hyper_parameter_tuning.tuner import Tuner
|
|
116
|
+
|
|
117
|
+
tuner = Tuner(
|
|
118
|
+
X=X_train, y=y_train,
|
|
119
|
+
model_name="lightgbm",
|
|
120
|
+
X_val=X_val, y_val=y_val,
|
|
121
|
+
X_test=X_test, y_test=y_test,
|
|
122
|
+
search_method="bayesian",
|
|
123
|
+
n_trials=50,
|
|
124
|
+
optimization_metric="prauc",
|
|
125
|
+
)
|
|
126
|
+
result = tuner.tune()
|
|
127
|
+
|
|
128
|
+
best_model = result["best_model"]
|
|
129
|
+
result["trials_results"] # full trial results (DataFrame)
|
|
130
|
+
```
|
|
131
|
+
|
|
132
|
+
Reported metrics are `auc`, `prauc`, `ks`, `log_loss` and `brier` (`ks`/`brier` are binary-only). Custom parameter search spaces can be defined via `ModelParams` and `ParamRange`.
|
|
133
|
+
|
|
134
|
+
- **Risk Class**: Builds risk classes from model scores. `RiskClassBuilderDP` searches bin edges with dynamic programming so that each class' observed event rate falls in a target band (`target_ranges`, required), then checks that adjacent classes are statistically distinguishable; `RiskClassBuilder` discovers classes with KMeans.
|
|
112
135
|
|
|
113
|
-
|
|
114
|
-
|
|
136
|
+
```python
|
|
137
|
+
from classifier_toolkit.risk_class import RiskClassBuilderDP
|
|
115
138
|
|
|
116
|
-
|
|
117
|
-
|
|
139
|
+
builder = RiskClassBuilderDP(
|
|
140
|
+
target_col="target",
|
|
141
|
+
target_ranges=[(0.00, 0.02), (0.02, 0.05), (0.05, 0.10)],
|
|
142
|
+
min_obs_per_bin=200,
|
|
143
|
+
)
|
|
144
|
+
result = builder.build(train_proba, y_train)
|
|
145
|
+
print(result.bins, result.n_classes)
|
|
118
146
|
```
|
|
119
147
|
|
|
120
|
-
|
|
148
|
+
- **Model Validation**: `Evaluator` evaluates any model with `predict_proba` (metrics, ROC/PR/calibration/threshold plots), `score_distribution_psi` and `compare_score_distributions` compare a reference score distribution with a current one, and `evaluate_risk_classes_dp` checks calibration within each risk class.
|
|
149
|
+
|
|
150
|
+
- **Calibration**: `OVRHistogramCalibrator` calibrates multiclass (or binary) probabilities One-vs-Rest with histogram binning; `plot_reliability_curves` shows raw vs calibrated reliability per class.
|
|
151
|
+
|
|
152
|
+
- **Explainability**: `TreeSHAPExplainer`, `ExplainabilityToolkit`, SHAP plots, pairwise interaction analysis and a `MisclassificationAnalyzer` that explains confusion-matrix quadrants (per-class SHAP for multiclass models).
|
|
153
|
+
|
|
154
|
+
- **Data Partition**: Temporal-aware train/test splitting, preprocessing helpers and dtype optimization.
|
|
121
155
|
|
|
122
156
|
### Development & CI/CD
|
|
123
157
|
|
|
@@ -131,8 +165,8 @@ This project uses modern tooling for fast and efficient development workflows:
|
|
|
131
165
|
#### CI/CD Pipeline
|
|
132
166
|
Our CI/CD pipeline is optimized for speed and efficiency:
|
|
133
167
|
|
|
134
|
-
- **Parallel Test Execution**:
|
|
135
|
-
- **Shared Caching**:
|
|
168
|
+
- **Parallel Test Execution**: One test job per directory under `tests/` (discovered automatically, so new test groups are picked up without editing the workflow), all running simultaneously
|
|
169
|
+
- **Shared Caching**: The parallel jobs share the same dependency cache, avoiding duplicate downloads
|
|
136
170
|
- **Smart Test Reruns**: Failed tests run first (`pytest --lf --ff`) for faster feedback on fixes
|
|
137
171
|
- **Master Protection**: Build tests only run on `master` branch and PRs targeting `master`, saving CI resources on feature branches
|
|
138
172
|
- **Automatic Linting**: Code quality checks (Ruff, SQLFluff) run on every push
|
|
@@ -157,5 +191,5 @@ uv build
|
|
|
157
191
|
|
|
158
192
|
### Future Work
|
|
159
193
|
The next planned improvements and additions to the library include:
|
|
160
|
-
*
|
|
194
|
+
* Extending the evaluation and reporting tools in `model_validation` (more checks and report formats).
|
|
161
195
|
* Expanding documentation to include architecture diagrams and detailed usage examples.
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
"""Probability calibration for multiclass (and binary) model output.
|
|
2
|
+
|
|
3
|
+
Public API for the OVR histogram-binning calibrator and its reliability
|
|
4
|
+
curve diagnostics — see :mod:`classifier_toolkit.calibration.ovr_calibration`
|
|
5
|
+
for details.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
import logging
|
|
9
|
+
from importlib import import_module
|
|
10
|
+
from typing import Any
|
|
11
|
+
|
|
12
|
+
logging.getLogger("classifier_toolkit.calibration").addHandler(logging.NullHandler())
|
|
13
|
+
|
|
14
|
+
_EXPORTS: dict[str, str] = {
|
|
15
|
+
"CalibrationError": "classifier_toolkit.calibration.base",
|
|
16
|
+
"OVRHistogramCalibrator": "classifier_toolkit.calibration.ovr_calibration",
|
|
17
|
+
"plot_reliability_curves": "classifier_toolkit.calibration.reliability",
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
__all__ = list(_EXPORTS.keys())
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def __getattr__(name: str) -> Any: # pragma: no cover
|
|
24
|
+
module_path = _EXPORTS.get(name)
|
|
25
|
+
if module_path is None:
|
|
26
|
+
raise AttributeError(
|
|
27
|
+
f"module 'classifier_toolkit.calibration' has no attribute '{name}'"
|
|
28
|
+
)
|
|
29
|
+
module = import_module(module_path)
|
|
30
|
+
return getattr(module, name)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def __dir__(): # pragma: no cover
|
|
34
|
+
return sorted(list(globals().keys()) + __all__)
|
|
@@ -0,0 +1,254 @@
|
|
|
1
|
+
"""One-vs-Rest histogram-binning calibration for multiclass probabilities.
|
|
2
|
+
|
|
3
|
+
Raw CatBoost/LightGBM scores are not guaranteed to be calibrated, which
|
|
4
|
+
matters whenever predicted probabilities are used as probabilities (pricing,
|
|
5
|
+
provisioning, or combining several models' outputs). This module calibrates
|
|
6
|
+
each class independently via histogram binning, then renormalizes the
|
|
7
|
+
calibrated row back to sum to 1.
|
|
8
|
+
|
|
9
|
+
It calibrates the output of a *single* model across the classes it predicts.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import warnings
|
|
13
|
+
from typing import Dict, Optional, Sequence, Union
|
|
14
|
+
|
|
15
|
+
import numpy as np
|
|
16
|
+
|
|
17
|
+
from classifier_toolkit.calibration.base import CalibrationError
|
|
18
|
+
|
|
19
|
+
ClassLabel = Union[str, int]
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class OVRHistogramCalibrator:
|
|
23
|
+
"""Calibrate multiclass probabilities via One-vs-Rest histogram binning.
|
|
24
|
+
|
|
25
|
+
For each class independently: bin the raw predicted probability into
|
|
26
|
+
``n_bins`` quantile-width bins, and replace it with the empirical class
|
|
27
|
+
frequency observed in that bin at fit time (histogram/binning
|
|
28
|
+
calibration — a non-parametric alternative to isotonic regression or
|
|
29
|
+
Platt scaling). Because each class is calibrated independently, the
|
|
30
|
+
calibrated row no longer sums to 1 — :meth:`transform` renormalizes it.
|
|
31
|
+
|
|
32
|
+
Sparse cells (e.g. very few observed events of a rare class in a given
|
|
33
|
+
probability bin) are shrunk toward that class's overall prevalence in
|
|
34
|
+
proportion to how little data the bin has, rather than reported as a
|
|
35
|
+
noisy raw frequency — see ``min_samples_per_bin``.
|
|
36
|
+
|
|
37
|
+
Parameters
|
|
38
|
+
----------
|
|
39
|
+
n_bins : int, optional
|
|
40
|
+
Target number of quantile bins per class, by default 10. When many
|
|
41
|
+
scores are tied (common for a rare class, e.g. most D scores at the
|
|
42
|
+
same low value), plain quantile edges collapse; each heavily tied
|
|
43
|
+
value then gets its own bin and the remaining bins are spread over
|
|
44
|
+
the untied scores. If fewer than 3 bins can still be formed (the
|
|
45
|
+
scores take very few distinct values), a warning is raised, since
|
|
46
|
+
the calibrated output for that class is then nearly constant.
|
|
47
|
+
min_samples_per_bin : int, optional
|
|
48
|
+
Minimum training samples a bin must contain before its empirical
|
|
49
|
+
frequency is trusted as-is; bins with fewer samples are shrunk
|
|
50
|
+
toward the class's overall prevalence in proportion to
|
|
51
|
+
``n_in_bin / min_samples_per_bin``. By default 20.
|
|
52
|
+
|
|
53
|
+
Attributes
|
|
54
|
+
----------
|
|
55
|
+
classes_ : Optional[np.ndarray]
|
|
56
|
+
Class labels, in the order matching probability columns. Set by fit.
|
|
57
|
+
bin_edges_ : Dict[ClassLabel, np.ndarray]
|
|
58
|
+
Per-class quantile bin edges from fit (outer edges are +/-inf so
|
|
59
|
+
every future probability falls inside some bin).
|
|
60
|
+
bin_values_ : Dict[ClassLabel, np.ndarray]
|
|
61
|
+
Per-class calibrated (possibly shrunk) frequency for each bin.
|
|
62
|
+
class_prior_ : Dict[ClassLabel, float]
|
|
63
|
+
Per-class overall prevalence at fit time, used for shrinkage.
|
|
64
|
+
"""
|
|
65
|
+
|
|
66
|
+
def __init__(self, n_bins: int = 10, min_samples_per_bin: int = 20) -> None:
|
|
67
|
+
self.n_bins = n_bins
|
|
68
|
+
self.min_samples_per_bin = min_samples_per_bin
|
|
69
|
+
self.classes_: Optional[np.ndarray] = None
|
|
70
|
+
self.bin_edges_: Dict[ClassLabel, np.ndarray] = {}
|
|
71
|
+
self.bin_values_: Dict[ClassLabel, np.ndarray] = {}
|
|
72
|
+
self.class_prior_: Dict[ClassLabel, float] = {}
|
|
73
|
+
|
|
74
|
+
def _quantile_edges(self, p: np.ndarray) -> np.ndarray:
|
|
75
|
+
quantiles = np.linspace(0.0, 1.0, self.n_bins + 1)
|
|
76
|
+
edges = np.unique(np.quantile(p, quantiles))
|
|
77
|
+
if len(edges) - 1 < self.n_bins:
|
|
78
|
+
edges = self._tie_aware_edges(p)
|
|
79
|
+
if len(edges) < 2:
|
|
80
|
+
edges = np.array([0.0, 1.0])
|
|
81
|
+
edges = edges.astype(float)
|
|
82
|
+
edges[0] = -np.inf
|
|
83
|
+
edges[-1] = np.inf
|
|
84
|
+
return edges
|
|
85
|
+
|
|
86
|
+
def _tie_aware_edges(self, p: np.ndarray) -> np.ndarray:
|
|
87
|
+
"""Quantile edges that give each heavily tied score its own bin.
|
|
88
|
+
|
|
89
|
+
A value is "heavily tied" when it holds more than ``1 / n_bins`` of
|
|
90
|
+
the samples, i.e. more than one quantile bin's worth. Such a value
|
|
91
|
+
collapses neighbouring quantile edges onto itself, so it is isolated
|
|
92
|
+
in its own bin (``[v, next float after v)``) and the rest of the bin
|
|
93
|
+
budget is spent on quantiles of the other scores.
|
|
94
|
+
"""
|
|
95
|
+
values, counts = np.unique(p, return_counts=True)
|
|
96
|
+
spikes = values[counts > len(p) / self.n_bins]
|
|
97
|
+
rest = p[~np.isin(p, spikes)]
|
|
98
|
+
n_rest_bins = max(self.n_bins - len(spikes), 1)
|
|
99
|
+
edges = list(spikes) + [np.nextafter(v, np.inf) for v in spikes]
|
|
100
|
+
if rest.size:
|
|
101
|
+
edges += list(np.quantile(rest, np.linspace(0.0, 1.0, n_rest_bins + 1)))
|
|
102
|
+
return np.unique(np.asarray(edges, dtype=float))
|
|
103
|
+
|
|
104
|
+
@staticmethod
|
|
105
|
+
def _bin_index(p: np.ndarray, edges: np.ndarray) -> np.ndarray:
|
|
106
|
+
return np.clip(np.digitize(p, edges[1:-1], right=False), 0, len(edges) - 2)
|
|
107
|
+
|
|
108
|
+
def fit(
|
|
109
|
+
self,
|
|
110
|
+
y_true: Union[Sequence, np.ndarray],
|
|
111
|
+
y_proba: np.ndarray,
|
|
112
|
+
classes: Optional[Sequence[ClassLabel]] = None,
|
|
113
|
+
) -> "OVRHistogramCalibrator":
|
|
114
|
+
"""Fit per-class histogram bins from raw predictions.
|
|
115
|
+
|
|
116
|
+
Parameters
|
|
117
|
+
----------
|
|
118
|
+
y_true : array-like of shape (n_samples,)
|
|
119
|
+
True class labels.
|
|
120
|
+
y_proba : np.ndarray of shape (n_samples, n_classes)
|
|
121
|
+
Raw (uncalibrated) predicted probabilities, e.g. from
|
|
122
|
+
``model.predict_proba(X)``.
|
|
123
|
+
classes : Sequence[ClassLabel], optional
|
|
124
|
+
Class labels in the order matching ``y_proba``'s columns
|
|
125
|
+
(typically ``model.classes_``). If None, inferred as the sorted
|
|
126
|
+
unique values of ``y_true`` — only safe if every class appears
|
|
127
|
+
in ``y_true`` and ``y_proba``'s columns are already in that
|
|
128
|
+
sorted order.
|
|
129
|
+
|
|
130
|
+
Returns
|
|
131
|
+
-------
|
|
132
|
+
OVRHistogramCalibrator
|
|
133
|
+
The fitted calibrator.
|
|
134
|
+
"""
|
|
135
|
+
y_true = np.asarray(y_true)
|
|
136
|
+
y_proba = self._check_proba(y_proba)
|
|
137
|
+
if classes is None:
|
|
138
|
+
classes = np.unique(y_true)
|
|
139
|
+
classes = np.asarray(classes)
|
|
140
|
+
if y_proba.shape[1] != len(classes):
|
|
141
|
+
raise CalibrationError(
|
|
142
|
+
f"y_proba has {y_proba.shape[1]} columns but {len(classes)} "
|
|
143
|
+
"classes were given."
|
|
144
|
+
)
|
|
145
|
+
if y_proba.shape[0] != len(y_true):
|
|
146
|
+
raise CalibrationError(
|
|
147
|
+
f"y_proba has {y_proba.shape[0]} rows but y_true has {len(y_true)}."
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
self.classes_ = classes
|
|
151
|
+
self.bin_edges_ = {}
|
|
152
|
+
self.bin_values_ = {}
|
|
153
|
+
self.class_prior_ = {}
|
|
154
|
+
|
|
155
|
+
for i, c in enumerate(classes):
|
|
156
|
+
y_bin = (y_true == c).astype(int)
|
|
157
|
+
p = y_proba[:, i]
|
|
158
|
+
prior = float(y_bin.mean()) if len(y_bin) > 0 else 0.0
|
|
159
|
+
self.class_prior_[c] = prior
|
|
160
|
+
|
|
161
|
+
edges = self._quantile_edges(p)
|
|
162
|
+
bin_idx = self._bin_index(p, edges)
|
|
163
|
+
n_populated = len(np.unique(bin_idx))
|
|
164
|
+
if n_populated < min(self.n_bins, 3):
|
|
165
|
+
warnings.warn(
|
|
166
|
+
f"Class {c!r}: only {n_populated} calibration bin(s) could "
|
|
167
|
+
f"be formed (n_bins={self.n_bins}) because its scores take "
|
|
168
|
+
"very few distinct values; its calibrated probability will "
|
|
169
|
+
"be nearly constant.",
|
|
170
|
+
stacklevel=2,
|
|
171
|
+
)
|
|
172
|
+
|
|
173
|
+
values = np.empty(len(edges) - 1, dtype=float)
|
|
174
|
+
for b in range(len(edges) - 1):
|
|
175
|
+
mask = bin_idx == b
|
|
176
|
+
n_in_bin = int(mask.sum())
|
|
177
|
+
if n_in_bin == 0:
|
|
178
|
+
values[b] = prior
|
|
179
|
+
continue
|
|
180
|
+
observed = float(y_bin[mask].mean())
|
|
181
|
+
if n_in_bin < self.min_samples_per_bin:
|
|
182
|
+
weight = n_in_bin / self.min_samples_per_bin
|
|
183
|
+
values[b] = weight * observed + (1 - weight) * prior
|
|
184
|
+
else:
|
|
185
|
+
values[b] = observed
|
|
186
|
+
self.bin_edges_[c] = edges
|
|
187
|
+
self.bin_values_[c] = values
|
|
188
|
+
|
|
189
|
+
return self
|
|
190
|
+
|
|
191
|
+
def transform(self, y_proba: np.ndarray) -> np.ndarray:
|
|
192
|
+
"""Calibrate raw probabilities and renormalize to sum to 1.
|
|
193
|
+
|
|
194
|
+
Parameters
|
|
195
|
+
----------
|
|
196
|
+
y_proba : np.ndarray of shape (n_samples, n_classes)
|
|
197
|
+
Raw predicted probabilities, columns matching ``classes_``.
|
|
198
|
+
|
|
199
|
+
Returns
|
|
200
|
+
-------
|
|
201
|
+
np.ndarray of shape (n_samples, n_classes)
|
|
202
|
+
Calibrated probabilities; each row sums to 1. A row whose
|
|
203
|
+
calibrated values are all exactly 0 (every relevant bin observed
|
|
204
|
+
no events) carries no information, so it falls back to the
|
|
205
|
+
class priors from fit.
|
|
206
|
+
"""
|
|
207
|
+
if self.classes_ is None:
|
|
208
|
+
raise CalibrationError("Calibrator is not fitted. Call fit() first.")
|
|
209
|
+
y_proba = self._check_proba(y_proba)
|
|
210
|
+
if y_proba.shape[1] != len(self.classes_):
|
|
211
|
+
raise CalibrationError(
|
|
212
|
+
f"y_proba has {y_proba.shape[1]} columns but the calibrator "
|
|
213
|
+
f"was fit on {len(self.classes_)} classes."
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
calibrated = np.empty_like(y_proba)
|
|
217
|
+
for i, c in enumerate(self.classes_):
|
|
218
|
+
edges = self.bin_edges_[c]
|
|
219
|
+
values = self.bin_values_[c]
|
|
220
|
+
bin_idx = self._bin_index(y_proba[:, i], edges)
|
|
221
|
+
calibrated[:, i] = values[bin_idx]
|
|
222
|
+
|
|
223
|
+
row_sums = calibrated.sum(axis=1, keepdims=True)
|
|
224
|
+
empty_rows = row_sums[:, 0] == 0
|
|
225
|
+
if empty_rows.any():
|
|
226
|
+
priors = np.array([self.class_prior_[c] for c in self.classes_])
|
|
227
|
+
if priors.sum() == 0:
|
|
228
|
+
priors = np.ones(len(self.classes_))
|
|
229
|
+
calibrated[empty_rows] = priors
|
|
230
|
+
row_sums = calibrated.sum(axis=1, keepdims=True)
|
|
231
|
+
return calibrated / row_sums
|
|
232
|
+
|
|
233
|
+
@staticmethod
|
|
234
|
+
def _check_proba(y_proba: np.ndarray) -> np.ndarray:
|
|
235
|
+
y_proba = np.asarray(y_proba, dtype=float)
|
|
236
|
+
if y_proba.ndim != 2:
|
|
237
|
+
raise CalibrationError(
|
|
238
|
+
f"y_proba must be 2D (n_samples, n_classes); got shape {y_proba.shape}."
|
|
239
|
+
)
|
|
240
|
+
if np.isnan(y_proba).any():
|
|
241
|
+
raise CalibrationError(
|
|
242
|
+
f"y_proba contains {int(np.isnan(y_proba).sum())} NaN value(s); "
|
|
243
|
+
"calibration needs a probability for every sample and class."
|
|
244
|
+
)
|
|
245
|
+
return y_proba
|
|
246
|
+
|
|
247
|
+
def fit_transform(
|
|
248
|
+
self,
|
|
249
|
+
y_true: Union[Sequence, np.ndarray],
|
|
250
|
+
y_proba: np.ndarray,
|
|
251
|
+
classes: Optional[Sequence[ClassLabel]] = None,
|
|
252
|
+
) -> np.ndarray:
|
|
253
|
+
"""Fit the calibrator and calibrate the same data in one step."""
|
|
254
|
+
return self.fit(y_true, y_proba, classes=classes).transform(y_proba)
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
"""Reliability curve diagnostics for OVR calibration.
|
|
2
|
+
|
|
3
|
+
One reliability (calibration) curve per class, each an
|
|
4
|
+
independent binary "class *c* vs rest" diagnostic — matching the OVR
|
|
5
|
+
calibration approach itself. Plots raw and (optionally) calibrated curves
|
|
6
|
+
together so the effect of calibration is visible directly.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import math
|
|
10
|
+
from typing import Optional, Sequence, Union
|
|
11
|
+
|
|
12
|
+
import matplotlib.pyplot as plt
|
|
13
|
+
import numpy as np
|
|
14
|
+
from sklearn.calibration import calibration_curve
|
|
15
|
+
|
|
16
|
+
ClassLabel = Union[str, int]
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def plot_reliability_curves(
|
|
20
|
+
y_true: Union[Sequence, np.ndarray],
|
|
21
|
+
y_proba_raw: np.ndarray,
|
|
22
|
+
y_proba_calibrated: Optional[np.ndarray] = None,
|
|
23
|
+
classes: Optional[Sequence[ClassLabel]] = None,
|
|
24
|
+
n_bins: int = 10,
|
|
25
|
+
pos_label: Optional[ClassLabel] = None,
|
|
26
|
+
n_cols: int = 3,
|
|
27
|
+
show: bool = True,
|
|
28
|
+
) -> plt.Figure:
|
|
29
|
+
"""Plot a reliability curve per class, one subplot each.
|
|
30
|
+
|
|
31
|
+
Parameters
|
|
32
|
+
----------
|
|
33
|
+
y_true : array-like of shape (n_samples,)
|
|
34
|
+
True class labels.
|
|
35
|
+
y_proba_raw : np.ndarray of shape (n_samples, n_classes)
|
|
36
|
+
Raw (uncalibrated) predicted probabilities.
|
|
37
|
+
y_proba_calibrated : np.ndarray of shape (n_samples, n_classes), optional
|
|
38
|
+
Calibrated predicted probabilities (e.g. from
|
|
39
|
+
:meth:`OVRHistogramCalibrator.transform`), plotted alongside the raw
|
|
40
|
+
curve for comparison. If None, only the raw curve is shown.
|
|
41
|
+
classes : Sequence[ClassLabel], optional
|
|
42
|
+
Class labels in the order matching the probability columns. If
|
|
43
|
+
None, inferred as the sorted unique values of ``y_true``.
|
|
44
|
+
n_bins : int, optional
|
|
45
|
+
Number of quantile bins per curve, by default 10.
|
|
46
|
+
pos_label : ClassLabel, optional
|
|
47
|
+
A class of interest (e.g. the default class). If given and present
|
|
48
|
+
in ``classes``, that subplot's title is bolded in red to keep it
|
|
49
|
+
visually prominent, consistent with the rest of the multiclass
|
|
50
|
+
toolkit. By default None.
|
|
51
|
+
n_cols : int, optional
|
|
52
|
+
Subplot grid columns, by default 3.
|
|
53
|
+
show : bool, optional
|
|
54
|
+
Whether to call ``plt.show()``, by default True.
|
|
55
|
+
|
|
56
|
+
Returns
|
|
57
|
+
-------
|
|
58
|
+
matplotlib.figure.Figure
|
|
59
|
+
"""
|
|
60
|
+
y_true = np.asarray(y_true)
|
|
61
|
+
y_proba_raw = np.asarray(y_proba_raw)
|
|
62
|
+
if classes is None:
|
|
63
|
+
classes = np.unique(y_true)
|
|
64
|
+
classes = list(classes)
|
|
65
|
+
n_classes = len(classes)
|
|
66
|
+
|
|
67
|
+
n_cols = min(n_cols, n_classes)
|
|
68
|
+
n_rows = math.ceil(n_classes / n_cols)
|
|
69
|
+
fig, axes = plt.subplots(
|
|
70
|
+
n_rows, n_cols, figsize=(5 * n_cols, 4 * n_rows), squeeze=False
|
|
71
|
+
)
|
|
72
|
+
axes_flat = axes.flatten()
|
|
73
|
+
|
|
74
|
+
for i, c in enumerate(classes):
|
|
75
|
+
ax = axes_flat[i]
|
|
76
|
+
y_bin = (y_true == c).astype(int)
|
|
77
|
+
|
|
78
|
+
frac_pos, mean_pred = calibration_curve(
|
|
79
|
+
y_bin, y_proba_raw[:, i], n_bins=n_bins, strategy="quantile"
|
|
80
|
+
)
|
|
81
|
+
ax.plot(mean_pred, frac_pos, marker="o", label="Raw")
|
|
82
|
+
|
|
83
|
+
if y_proba_calibrated is not None:
|
|
84
|
+
frac_pos_cal, mean_pred_cal = calibration_curve(
|
|
85
|
+
y_bin,
|
|
86
|
+
np.asarray(y_proba_calibrated)[:, i],
|
|
87
|
+
n_bins=n_bins,
|
|
88
|
+
strategy="quantile",
|
|
89
|
+
)
|
|
90
|
+
ax.plot(mean_pred_cal, frac_pos_cal, marker="s", label="Calibrated")
|
|
91
|
+
|
|
92
|
+
ax.plot([0, 1], [0, 1], linestyle="--", color="gray", label="Perfect")
|
|
93
|
+
ax.set_xlabel("Mean predicted probability")
|
|
94
|
+
ax.set_ylabel("Observed frequency")
|
|
95
|
+
title = f"Class {c}"
|
|
96
|
+
if pos_label is not None and c == pos_label:
|
|
97
|
+
ax.set_title(title, fontweight="bold", color="red")
|
|
98
|
+
else:
|
|
99
|
+
ax.set_title(title)
|
|
100
|
+
ax.legend(fontsize=8)
|
|
101
|
+
|
|
102
|
+
for i in range(n_classes, n_rows * n_cols):
|
|
103
|
+
axes_flat[i].set_visible(False)
|
|
104
|
+
|
|
105
|
+
fig.suptitle("Reliability Curves (One-vs-Rest per class)", fontsize=14, y=1.02)
|
|
106
|
+
plt.tight_layout()
|
|
107
|
+
if show:
|
|
108
|
+
plt.show()
|
|
109
|
+
return fig
|