causal-testing-framework 14.2.0__tar.gz → 14.3.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.gitignore +2 -0
- {causal_testing_framework-14.2.0/causal_testing_framework.egg-info → causal_testing_framework-14.3.0}/PKG-INFO +2 -1
- causal_testing_framework-14.3.0/causal_testing/__main__.py +137 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/_version.py +3 -3
- causal_testing_framework-14.3.0/causal_testing/discovery/abstract_discovery.py +236 -0
- causal_testing_framework-14.3.0/causal_testing/discovery/hill_climber_discovery.py +155 -0
- causal_testing_framework-14.3.0/causal_testing/discovery/nsga_discovery.py +131 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/abstract_regression_estimator.py +15 -1
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/linear_regression_estimator.py +4 -4
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/logistic_regression_estimator.py +2 -8
- causal_testing_framework-14.3.0/causal_testing/estimation/multinomial_regression_estimator.py +56 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/main.py +63 -4
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/specification/causal_dag.py +5 -4
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/causal_effect.py +2 -2
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/metamorphic_relation.py +28 -11
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0/causal_testing_framework.egg-info}/PKG-INFO +2 -1
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing_framework.egg-info/SOURCES.txt +14 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing_framework.egg-info/entry_points.txt +5 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing_framework.egg-info/requires.txt +1 -0
- causal_testing_framework-14.3.0/causal_testing_framework.egg-info/scm_file_list.json +188 -0
- causal_testing_framework-14.3.0/causal_testing_framework.egg-info/scm_version.json +8 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/index.rst +1 -0
- causal_testing_framework-14.3.0/docs/source/modules/discovery.rst +82 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/modules/estimators.rst +12 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/vaccinating_elderly/causal_test_results.json +3 -3
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/pyproject.toml +21 -15
- causal_testing_framework-14.3.0/tests/__init__.py +0 -0
- causal_testing_framework-14.3.0/tests/discovery_tests/test_abstract_discovery.py +262 -0
- causal_testing_framework-14.3.0/tests/discovery_tests/test_hill_climber_discovery.py +132 -0
- causal_testing_framework-14.3.0/tests/discovery_tests/test_nsga_discovery.py +53 -0
- causal_testing_framework-14.3.0/tests/estimation_tests/test_multinomial_regression_estimator.py +42 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/main_tests/test_main.py +59 -1
- causal_testing_framework-14.3.0/tests/resources/data/exclude_edges.dot +3 -0
- causal_testing_framework-14.3.0/tests/resources/data/include_edges.dot +3 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/testing_tests/test_metamorphic_relations.py +29 -62
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/tutorial_tests/test_tutorials.py +1 -1
- causal_testing_framework-14.2.0/causal_testing/__main__.py +0 -94
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/ISSUE_TEMPLATE/bug_report.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/ISSUE_TEMPLATE/feature_request.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/ci-tests-drafts.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/ci-tests.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/figshare.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/joss.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/lint-format.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/publish-to-dafni.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.github/workflows/publish-to-pypi.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.mega-linter.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.pre-commit-config.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.pylintrc +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/.readthedocs.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/CITATION.cff +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/CONTRIBUTING.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/LICENSE +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/__init__.py +0 -0
- {causal_testing_framework-14.2.0/causal_testing/estimation → causal_testing_framework-14.3.0/causal_testing/discovery}/__init__.py +0 -0
- {causal_testing_framework-14.2.0/causal_testing/specification → causal_testing_framework-14.3.0/causal_testing/estimation}/__init__.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/abstract_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/cubic_spline_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/effect_estimate.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/experimental_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/genetic_programming_regression_fitter.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/instrumental_variable_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/estimation/ipcw_estimator.py +0 -0
- {causal_testing_framework-14.2.0/causal_testing/surrogate → causal_testing_framework-14.3.0/causal_testing/specification}/__init__.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/specification/scenario.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/specification/variable.py +0 -0
- {causal_testing_framework-14.2.0/causal_testing/testing → causal_testing_framework-14.3.0/causal_testing/surrogate}/__init__.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/surrogate/causal_surrogate_assisted.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/surrogate/surrogate_search_algorithms.py +0 -0
- {causal_testing_framework-14.2.0/causal_testing/utils → causal_testing_framework-14.3.0/causal_testing/testing}/__init__.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/base_test_case.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/causal_test_adequacy.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/causal_test_case.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/causal_test_result.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/effect.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/testing/intervention.py +0 -0
- {causal_testing_framework-14.2.0/tests → causal_testing_framework-14.3.0/causal_testing/utils}/__init__.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/utils/validation.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing_framework.egg-info/dependency_links.txt +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing_framework.egg-info/top_level.txt +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/codecov.yml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/.dockerignore +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/.env +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/Dockerfile +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/data/inputs/causal-tests/causal_tests.json +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/data/inputs/dag-data/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/data/inputs/runtime-data/runtime_data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/data/outputs/causal_test_results.json +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/docker-compose.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/entrypoint.sh +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/dafni/model_definition.yaml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/Makefile +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/make.bat +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/_static/css/custom.css +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/_static/images/CITCOM-logo-white.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/_static/images/CITCOM-logo.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/_static/images/Sheffield-logo.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/_static/images/example_dag.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/background.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/conf.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/credits.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/dev/actions_and_webhooks.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/dev/documentation.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/dev/version_release.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/glossary.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/installation.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/modules/causal_specification.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/modules/causal_testing.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/modules/custom_estimators.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/requirements.txt +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/poisson_line_process/causal_tests.json +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/poisson_line_process/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/poisson_line_process/dag.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/poisson_line_process/data/random/data_random_1000.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/poisson_line_process/poisson_line_process_tutorial.ipynb +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/vaccinating_elderly/causal_tests.json +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/vaccinating_elderly/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/vaccinating_elderly/dag_image.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/vaccinating_elderly/simulated_data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/vaccinating_elderly/vaccinating_elderly_tutorial.ipynb +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials/visualising_causal_test_results/visualise_causal_test_results.ipynb +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/docs/source/tutorials.rst +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/.gitignore +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/dag.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/10k_observational_data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/high_contacts_avg_age_22.2.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/high_contacts_avg_age_30.1.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/low_contacts_avg_age_22.2.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/low_contacts_avg_age_30.1.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/older_population.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/data/younger_population.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/doubling_beta/example_beta.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/vaccinating_elderly/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/vaccinating_elderly/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/vaccinating_elderly/dag.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/vaccinating_elderly/example_vaccine.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/covasim_/vaccinating_elderly/simulated_data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/lr91/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/lr91/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/lr91/dag.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/lr91/data/normalised_results.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/lr91/data/results.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/lr91/example_max_conductances.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/.gitignore +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/README.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/causal_tests.json +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/dag.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/random/data_random_1000.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh10_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh1_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh2_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh3_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh4_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh5_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh6_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh7_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh8_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/data/smt_100/data_smt_wh9_100.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/example_pure_python.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/examples/poisson-line-process/poisson_line_process.ipynb +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/images/.gitignore +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/images/schematic-dark.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/images/schematic.png +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/images/schematic.tex +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/paper/paper.bib +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/paper/paper.md +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/setup.cfg +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_cubic_spline_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_experimental_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_genetic_programming_regression_fitter.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_instrumental_variable_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_ipcw_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_linear_regression_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/estimation_tests/test_logistic_regression_estimator.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/dag.dot +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/dag.xml +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/data.pqt +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/data_with_categorical.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/data_with_meta.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/nhefs.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/scarf_data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/temporal_data.csv +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/resources/data/tests.json +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/specification_tests/test_causal_dag.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/specification_tests/test_variable.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/surrogate_tests/test_causal_surrogate_assisted.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/testing_tests/test_causal_effect.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/testing_tests/test_causal_test_adequacy.py +0 -0
- {causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/tests/testing_tests/test_causal_test_case.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: causal_testing_framework
|
|
3
|
-
Version: 14.
|
|
3
|
+
Version: 14.3.0
|
|
4
4
|
Summary: A framework for causal testing using causal directed acyclic graphs.
|
|
5
5
|
Author: The CITCOM team
|
|
6
6
|
License: MIT
|
|
@@ -27,6 +27,7 @@ Requires-Dist: sympy~=1.14.0
|
|
|
27
27
|
Requires-Dist: pyarrow<25,>=19.0.1
|
|
28
28
|
Requires-Dist: fastparquet>=2024.11.0
|
|
29
29
|
Requires-Dist: tqdm~=4.67.1
|
|
30
|
+
Requires-Dist: rustworkx>=0.18.0
|
|
30
31
|
Provides-Extra: dev
|
|
31
32
|
Requires-Dist: astroid==3.3.8; extra == "dev"
|
|
32
33
|
Requires-Dist: isort; extra == "dev"
|
|
@@ -0,0 +1,137 @@
|
|
|
1
|
+
"""This module contains the main entrypoint functionality to the Causal Testing Framework."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import logging
|
|
5
|
+
import os
|
|
6
|
+
import tempfile
|
|
7
|
+
from importlib.metadata import entry_points
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
|
|
10
|
+
import networkx as nx
|
|
11
|
+
import pandas as pd
|
|
12
|
+
|
|
13
|
+
from causal_testing.testing.metamorphic_relation import generate_causal_tests
|
|
14
|
+
|
|
15
|
+
from .main import CausalTestingFramework, CausalTestingPaths, Command, parse_args, setup_logging
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def main() -> None:
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
Main entry point for the Causal Testing Framework
|
|
22
|
+
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
# Parse arguments
|
|
26
|
+
args = parse_args()
|
|
27
|
+
# Setup logging
|
|
28
|
+
setup_logging(args.log_level)
|
|
29
|
+
|
|
30
|
+
match args.command:
|
|
31
|
+
case Command.GENERATE:
|
|
32
|
+
logging.info("Generating causal tests")
|
|
33
|
+
generate_causal_tests(
|
|
34
|
+
args.dag_path,
|
|
35
|
+
args.output,
|
|
36
|
+
args.ignore_cycles,
|
|
37
|
+
args.threads,
|
|
38
|
+
effect_type=args.effect_type,
|
|
39
|
+
estimate_type=args.estimate_type,
|
|
40
|
+
estimator=args.estimator,
|
|
41
|
+
skip=False,
|
|
42
|
+
)
|
|
43
|
+
logging.info("Causal test generation completed successfully.")
|
|
44
|
+
|
|
45
|
+
case Command.DISCOVER:
|
|
46
|
+
discover_map = {ff.name: ff for ff in entry_points(group="discovery")}
|
|
47
|
+
if args.technique not in discover_map:
|
|
48
|
+
raise ValueError(
|
|
49
|
+
f"Unsupported technique {args.technique}. Supported: {sorted(discover_map)}. "
|
|
50
|
+
"If you have implemented a custom technique, you will need to add this to your entrypoints via "
|
|
51
|
+
"your pyproject.toml file."
|
|
52
|
+
)
|
|
53
|
+
kwargs = {}
|
|
54
|
+
for argument in args.technique_kwargs:
|
|
55
|
+
split = argument.split("=")
|
|
56
|
+
if len(split) != 2:
|
|
57
|
+
raise ValueError(f"Malformed argument {argument}. Should be specified as `arg_name=arg_value`")
|
|
58
|
+
kwargs[split[0]] = split[1]
|
|
59
|
+
|
|
60
|
+
logging.info("Discovering causal structure")
|
|
61
|
+
# Need to reset index to allow for multiple files having the same index (i.e. starting at zero).
|
|
62
|
+
# Otherwise you end up with duplicate indices, which causes problems further down the line
|
|
63
|
+
df = pd.concat([pd.read_csv(path) for path in args.data_paths]).reset_index()
|
|
64
|
+
if args.variables:
|
|
65
|
+
df = df[args.variables]
|
|
66
|
+
|
|
67
|
+
discover_class = discover_map[args.technique].load()
|
|
68
|
+
discover = discover_class(
|
|
69
|
+
df=df,
|
|
70
|
+
exclude_edges=(
|
|
71
|
+
list(nx.nx_pydot.read_dot(args.exclude_edges).edges()) if args.exclude_edges is not None else []
|
|
72
|
+
),
|
|
73
|
+
include_edges=(
|
|
74
|
+
list(nx.nx_pydot.read_dot(args.include_edges).edges()) if args.include_edges is not None else []
|
|
75
|
+
),
|
|
76
|
+
alpha=args.alpha,
|
|
77
|
+
**kwargs,
|
|
78
|
+
)
|
|
79
|
+
evolved_dag = discover.discover()
|
|
80
|
+
discover.write_dot(evolved_dag, args.output)
|
|
81
|
+
logging.info("Causal structure discovery completed successfully.")
|
|
82
|
+
case Command.TEST:
|
|
83
|
+
# Create paths object
|
|
84
|
+
paths = CausalTestingPaths(
|
|
85
|
+
dag_path=args.dag_path,
|
|
86
|
+
data_paths=args.data_paths,
|
|
87
|
+
test_config_path=args.test_config,
|
|
88
|
+
output_path=args.output,
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
# Create and setup framework
|
|
92
|
+
framework = CausalTestingFramework(paths, ignore_cycles=args.ignore_cycles, query=args.query)
|
|
93
|
+
framework.setup()
|
|
94
|
+
|
|
95
|
+
# Load and run tests
|
|
96
|
+
framework.load_tests()
|
|
97
|
+
|
|
98
|
+
if args.batch_size > 0:
|
|
99
|
+
logging.info(f"Running tests in batches of size {args.batch_size}")
|
|
100
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
101
|
+
output_files = []
|
|
102
|
+
for i, results in enumerate(
|
|
103
|
+
framework.run_tests_in_batches(
|
|
104
|
+
batch_size=args.batch_size,
|
|
105
|
+
silent=args.silent,
|
|
106
|
+
adequacy=args.adequacy,
|
|
107
|
+
bootstrap_size=args.bootstrap_size,
|
|
108
|
+
)
|
|
109
|
+
):
|
|
110
|
+
temp_file_path = os.path.join(tmpdir, f"output_{i}.json")
|
|
111
|
+
framework.save_results(results, temp_file_path)
|
|
112
|
+
output_files.append(temp_file_path)
|
|
113
|
+
del results
|
|
114
|
+
|
|
115
|
+
# Now stitch the results together from the temporary files
|
|
116
|
+
all_results = []
|
|
117
|
+
for file_path in output_files:
|
|
118
|
+
with open(file_path, "r", encoding="utf-8") as f:
|
|
119
|
+
all_results.extend(json.load(f))
|
|
120
|
+
|
|
121
|
+
output_path = Path(args.output)
|
|
122
|
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
|
123
|
+
|
|
124
|
+
with open(args.output, "w", encoding="utf-8") as f:
|
|
125
|
+
json.dump(all_results, f, indent=4)
|
|
126
|
+
else:
|
|
127
|
+
logging.info("Running tests in regular mode")
|
|
128
|
+
results = framework.run_tests(
|
|
129
|
+
silent=args.silent, adequacy=args.adequacy, bootstrap_size=args.bootstrap_size
|
|
130
|
+
)
|
|
131
|
+
framework.save_results(results)
|
|
132
|
+
|
|
133
|
+
logging.info("Causal testing completed successfully.")
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
if __name__ == "__main__":
|
|
137
|
+
main()
|
{causal_testing_framework-14.2.0 → causal_testing_framework-14.3.0}/causal_testing/_version.py
RENAMED
|
@@ -18,7 +18,7 @@ version_tuple: tuple[int | str, ...]
|
|
|
18
18
|
commit_id: str | None
|
|
19
19
|
__commit_id__: str | None
|
|
20
20
|
|
|
21
|
-
__version__ = version = '14.
|
|
22
|
-
__version_tuple__ = version_tuple = (14,
|
|
21
|
+
__version__ = version = '14.3.0'
|
|
22
|
+
__version_tuple__ = version_tuple = (14, 3, 0)
|
|
23
23
|
|
|
24
|
-
__commit_id__ = commit_id = '
|
|
24
|
+
__commit_id__ = commit_id = 'ge6a19c77b'
|
|
@@ -0,0 +1,236 @@
|
|
|
1
|
+
"""
|
|
2
|
+
This module implements the abstract Discovery class to infer causal DAGs from data.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import random
|
|
6
|
+
import re
|
|
7
|
+
import warnings
|
|
8
|
+
from abc import ABC, abstractmethod
|
|
9
|
+
from enum import Enum
|
|
10
|
+
from itertools import permutations
|
|
11
|
+
|
|
12
|
+
import networkx as nx
|
|
13
|
+
import pandas as pd
|
|
14
|
+
import rustworkx as rx
|
|
15
|
+
|
|
16
|
+
from causal_testing.main import CausalTestingFramework
|
|
17
|
+
from causal_testing.specification.causal_dag import CausalDAG
|
|
18
|
+
from causal_testing.specification.scenario import Scenario
|
|
19
|
+
from causal_testing.testing.causal_effect import Negative, Positive
|
|
20
|
+
from causal_testing.testing.causal_test_result import CausalTestResult
|
|
21
|
+
from causal_testing.testing.metamorphic_relation import generate_metamorphic_relations
|
|
22
|
+
|
|
23
|
+
TestResult = Enum("TestResult", [("PASS", 2), ("FAIL", 0), ("INESTIMABLE", 1)])
|
|
24
|
+
|
|
25
|
+
# Ignore warnings from statsmodels when we try to evaluate test cases
|
|
26
|
+
warnings.simplefilter("ignore")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def simple_cycle(causal_dag: CausalDAG):
|
|
30
|
+
"""
|
|
31
|
+
Find a cycle in the given CausalDAG, if one exists, returns the first found.
|
|
32
|
+
|
|
33
|
+
:param causal_dag: The CausalDAG to check for cycles.
|
|
34
|
+
:returns: A list of edges in the cycle, or an empty list if there are no cycles.
|
|
35
|
+
"""
|
|
36
|
+
rx_graph = rx.networkx_converter(causal_dag)
|
|
37
|
+
return [(rx_graph[i], rx_graph[j]) for i, j in rx.digraph_find_cycle(rx_graph)]
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def effect_direction(result: CausalTestResult) -> str:
|
|
41
|
+
"""
|
|
42
|
+
Check whether the estimated causal effect is negative or positive.
|
|
43
|
+
|
|
44
|
+
:param result: The causal test result object.
|
|
45
|
+
:returns: Whether the estimated causal test is positive or negative (or no effect).
|
|
46
|
+
"""
|
|
47
|
+
if pd.api.types.is_numeric_dtype(
|
|
48
|
+
result.estimator.df[result.estimator.base_test_case.treatment_variable.name]
|
|
49
|
+
) and pd.api.types.is_numeric_dtype(result.estimator.df[result.estimator.base_test_case.outcome_variable.name]):
|
|
50
|
+
if Negative().apply(result):
|
|
51
|
+
return "negative"
|
|
52
|
+
if Positive().apply(result):
|
|
53
|
+
return "positive"
|
|
54
|
+
return None
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def is_match(u: str, v: str, patterns: list[str]):
|
|
58
|
+
"""
|
|
59
|
+
Check whether a given edge matches a given pattern.
|
|
60
|
+
|
|
61
|
+
:param u: The origin node of the edge.
|
|
62
|
+
:param v: The destination node of the edge.
|
|
63
|
+
:param patterns: A list of tuples containing the patterns to check against.
|
|
64
|
+
:returns: True if the edge matches the pattern, False otherwise.
|
|
65
|
+
"""
|
|
66
|
+
return any(re.fullmatch(pat_u, u) and re.fullmatch(pat_v, v) for pat_u, pat_v in patterns)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
class Discovery(ABC):
|
|
70
|
+
"""
|
|
71
|
+
Abstract class for causal discovery.
|
|
72
|
+
"""
|
|
73
|
+
|
|
74
|
+
def __init__(
|
|
75
|
+
self,
|
|
76
|
+
df: pd.DataFrame,
|
|
77
|
+
random_seed: int = 0,
|
|
78
|
+
exclude_edges: str = None,
|
|
79
|
+
include_edges: str = None,
|
|
80
|
+
alpha: float = 0.05,
|
|
81
|
+
):
|
|
82
|
+
|
|
83
|
+
random.seed(random_seed)
|
|
84
|
+
self.df = df
|
|
85
|
+
self.random_seed = int(random_seed)
|
|
86
|
+
self.alpha = float(alpha)
|
|
87
|
+
|
|
88
|
+
self.possible_edges = []
|
|
89
|
+
self.include_edges = []
|
|
90
|
+
self.exclude_edges = []
|
|
91
|
+
|
|
92
|
+
for u, v in permutations(df.columns, 2):
|
|
93
|
+
if exclude_edges and is_match(u, v, exclude_edges):
|
|
94
|
+
self.exclude_edges.append((u, v))
|
|
95
|
+
else:
|
|
96
|
+
self.possible_edges.append((u, v))
|
|
97
|
+
|
|
98
|
+
if include_edges and is_match(u, v, include_edges):
|
|
99
|
+
self.include_edges.append((u, v))
|
|
100
|
+
|
|
101
|
+
if self.include_edges:
|
|
102
|
+
# Check to make sure that the include edges don't specify a cycle
|
|
103
|
+
initial_dag = CausalDAG()
|
|
104
|
+
initial_dag.add_edges_from(self.include_edges)
|
|
105
|
+
|
|
106
|
+
if not initial_dag.is_acyclic():
|
|
107
|
+
raise ValueError(
|
|
108
|
+
"Specified include edges include a cycle, making it impossible to infer a DAG. "
|
|
109
|
+
"Please resolve this and try again."
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
@abstractmethod
|
|
113
|
+
def discover(self) -> CausalDAG:
|
|
114
|
+
"""
|
|
115
|
+
Discover the causal DAG.
|
|
116
|
+
|
|
117
|
+
:returns: The inferred causal DAG.
|
|
118
|
+
"""
|
|
119
|
+
|
|
120
|
+
def remove_cycles(self, causal_dag: CausalDAG):
|
|
121
|
+
"""
|
|
122
|
+
Remove cycles from individuals by iteratively deleting a random edge from each cycle until there are no more
|
|
123
|
+
cycles.
|
|
124
|
+
|
|
125
|
+
:param causal_dag: The CausalDAG to be repaired.
|
|
126
|
+
"""
|
|
127
|
+
nodes = causal_dag.nodes
|
|
128
|
+
cycle = simple_cycle(causal_dag)
|
|
129
|
+
while cycle:
|
|
130
|
+
idx = random.choice(range(len(cycle)))
|
|
131
|
+
while cycle[idx] in self.include_edges:
|
|
132
|
+
idx = (idx + 1) % len(cycle)
|
|
133
|
+
causal_dag.remove_edge(cycle[idx][0], cycle[idx][1])
|
|
134
|
+
cycle = simple_cycle(causal_dag)
|
|
135
|
+
causal_dag.add_nodes_from(nodes)
|
|
136
|
+
|
|
137
|
+
def write_dot(self, individual: CausalDAG, output_file: str):
|
|
138
|
+
"""
|
|
139
|
+
Write the given individual to the given output file.
|
|
140
|
+
|
|
141
|
+
:param individual: The causal DAG to output.
|
|
142
|
+
:param output_file: The name of the file to write to.
|
|
143
|
+
"""
|
|
144
|
+
if hasattr(individual, "test_results"):
|
|
145
|
+
for _, test in individual.test_results.iterrows():
|
|
146
|
+
if (test["treatment"], test["outcome"]) in individual.edges:
|
|
147
|
+
if test["result"] == TestResult.PASS:
|
|
148
|
+
individual[test["treatment"]][test["outcome"]]["color"] = "green"
|
|
149
|
+
elif test["result"] == TestResult.INESTIMABLE:
|
|
150
|
+
individual[test["treatment"]][test["outcome"]]["color"] = "orange"
|
|
151
|
+
elif test["result"] == TestResult.FAIL:
|
|
152
|
+
individual[test["treatment"]][test["outcome"]]["color"] = "red"
|
|
153
|
+
else:
|
|
154
|
+
raise ValueError(f"Invalid test outcome {test['result']}")
|
|
155
|
+
else:
|
|
156
|
+
individual.add_edge(test["treatment"], test["outcome"], ignore_cycles=True)
|
|
157
|
+
individual[test["treatment"]][test["outcome"]]["style"] = "dashed"
|
|
158
|
+
if test["result"] == TestResult.PASS:
|
|
159
|
+
individual[test["treatment"]][test["outcome"]]["style"] = "invis"
|
|
160
|
+
individual[test["treatment"]][test["outcome"]]["constraint"] = False
|
|
161
|
+
elif test["result"] == TestResult.INESTIMABLE:
|
|
162
|
+
individual[test["treatment"]][test["outcome"]]["color"] = "orange"
|
|
163
|
+
elif test["result"] == TestResult.FAIL:
|
|
164
|
+
individual[test["treatment"]][test["outcome"]]["color"] = "red"
|
|
165
|
+
else:
|
|
166
|
+
raise ValueError(f"Invalid test outcome {test['result']}")
|
|
167
|
+
|
|
168
|
+
nx.drawing.nx_pydot.write_dot(individual, output_file)
|
|
169
|
+
|
|
170
|
+
def _json_stub_params(self, outcome: str) -> str:
|
|
171
|
+
if pd.api.types.is_bool_dtype(self.df[outcome]):
|
|
172
|
+
return {"estimator": "LogisticRegressionEstimator", "estimate_type": "unit_odds_ratio"}
|
|
173
|
+
if pd.api.types.is_categorical_dtype(self.df[outcome]) or pd.api.types.is_object_dtype(self.df[outcome]):
|
|
174
|
+
return {"estimator": "MultinomialRegressionEstimator", "estimate_type": "unit_odds_ratio"}
|
|
175
|
+
if pd.api.types.is_numeric_dtype(self.df[outcome]):
|
|
176
|
+
return {"estimator": "LinearRegressionEstimator", "estimate_type": "coefficient"}
|
|
177
|
+
raise ValueError(f"Invalid datatype {self.df.dtypes[outcome]}")
|
|
178
|
+
|
|
179
|
+
def evaluate_tests(self, causal_dag: CausalDAG) -> pd.DataFrame:
|
|
180
|
+
"""
|
|
181
|
+
Generate and evaluate causal test cases from the supplied CausalDAG and return a list of edges for which the
|
|
182
|
+
corresponding causal test case failed.
|
|
183
|
+
These results are then assigned to a new attribute `test_results` within the individual for later reuse.
|
|
184
|
+
|
|
185
|
+
:param causal_dag: The CausalDAG to evaluate.
|
|
186
|
+
:returns: Pandas dataframe with test outcome details
|
|
187
|
+
(result, expected effect, treatment, outcome, effect direction).
|
|
188
|
+
"""
|
|
189
|
+
|
|
190
|
+
ctf = CausalTestingFramework(None)
|
|
191
|
+
ctf.dag = causal_dag
|
|
192
|
+
ctf.data = self.df
|
|
193
|
+
ctf.create_variables()
|
|
194
|
+
ctf.scenario = Scenario(list(ctf.variables["inputs"].values()) + list(ctf.variables["outputs"].values()))
|
|
195
|
+
|
|
196
|
+
ctf.test_cases = ctf.create_test_cases(
|
|
197
|
+
{
|
|
198
|
+
"tests": [
|
|
199
|
+
relation.to_json_stub(
|
|
200
|
+
alpha=self.alpha,
|
|
201
|
+
**self._json_stub_params(relation.base_test_case.outcome_variable),
|
|
202
|
+
)
|
|
203
|
+
for relation in generate_metamorphic_relations(causal_dag)
|
|
204
|
+
]
|
|
205
|
+
}
|
|
206
|
+
)
|
|
207
|
+
|
|
208
|
+
results = []
|
|
209
|
+
|
|
210
|
+
# We use "silent=False" here to allow for inestimable edges, but it'd be good to have a more stringent
|
|
211
|
+
# error catching strategy to catch "genuine" problems (e.g. to do with the structure of the data)
|
|
212
|
+
for test_case, result in zip(ctf.test_cases, ctf.run_tests(silent=False)):
|
|
213
|
+
if result.effect_estimate is None:
|
|
214
|
+
results.append(
|
|
215
|
+
{
|
|
216
|
+
"result": TestResult.INESTIMABLE,
|
|
217
|
+
"expected_effect": test_case.expected_causal_effect.__class__.__name__,
|
|
218
|
+
"treatment": test_case.base_test_case.treatment_variable.name,
|
|
219
|
+
"outcome": test_case.base_test_case.outcome_variable.name,
|
|
220
|
+
}
|
|
221
|
+
)
|
|
222
|
+
else:
|
|
223
|
+
results.append(
|
|
224
|
+
{
|
|
225
|
+
"result": (
|
|
226
|
+
TestResult.PASS if test_case.expected_causal_effect.apply(result) else TestResult.FAIL
|
|
227
|
+
),
|
|
228
|
+
"expected_effect": test_case.expected_causal_effect.__class__.__name__,
|
|
229
|
+
"treatment": test_case.base_test_case.treatment_variable.name,
|
|
230
|
+
"outcome": test_case.base_test_case.outcome_variable.name,
|
|
231
|
+
"effect": effect_direction(result),
|
|
232
|
+
}
|
|
233
|
+
)
|
|
234
|
+
|
|
235
|
+
causal_dag.test_results = pd.DataFrame(results)
|
|
236
|
+
return pd.DataFrame(results)
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""
|
|
2
|
+
This module implements a hill climbing algorithm to optimise causal DAGs based on the tests that pass/fail.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import random
|
|
6
|
+
import time
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import pandas as pd
|
|
10
|
+
|
|
11
|
+
from causal_testing.discovery.abstract_discovery import Discovery, TestResult
|
|
12
|
+
from causal_testing.specification.causal_dag import CausalDAG
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class HillClimberDiscovery(Discovery):
|
|
16
|
+
"""
|
|
17
|
+
Simple hill climber evolution of cauasl DAGs via 1+1EA.
|
|
18
|
+
Attempts to maximise the number of passing tests and minimise the number of failing tests.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
def __init__(
|
|
22
|
+
self,
|
|
23
|
+
df: pd.DataFrame,
|
|
24
|
+
random_seed: int = 0,
|
|
25
|
+
include_edges: str = None,
|
|
26
|
+
exclude_edges: str = None,
|
|
27
|
+
alpha: float = 0.05,
|
|
28
|
+
max_iterations: int = 100,
|
|
29
|
+
max_iterations_without_improvement: int = 10,
|
|
30
|
+
):
|
|
31
|
+
super().__init__(
|
|
32
|
+
df=df, random_seed=random_seed, include_edges=include_edges, exclude_edges=exclude_edges, alpha=alpha
|
|
33
|
+
)
|
|
34
|
+
self.max_iterations = int(max_iterations)
|
|
35
|
+
self.max_iterations_without_improvement = int(max_iterations_without_improvement)
|
|
36
|
+
|
|
37
|
+
def sum_test_outcomes(self, test_results: pd.DataFrame) -> dict:
|
|
38
|
+
"""
|
|
39
|
+
Aggregate the number of passing, failing, and inestimable tests
|
|
40
|
+
:param test_results: Dataframe containing the raw pass/fail/inestimable outcome of each test case.
|
|
41
|
+
:returns: Dictionary containing the number of pass/fail/inestimable outcomes.
|
|
42
|
+
"""
|
|
43
|
+
counts = pd.concat(
|
|
44
|
+
[
|
|
45
|
+
pd.DataFrame(np.sort(test_results[["treatment", "outcome"]], axis=1), columns=["treatment", "outcome"]),
|
|
46
|
+
pd.get_dummies(test_results["result"]).astype(int),
|
|
47
|
+
],
|
|
48
|
+
axis=1,
|
|
49
|
+
)
|
|
50
|
+
# Ensure every column is initialised - Test outcomes that never occurred won't be in the dataframe otherwise
|
|
51
|
+
for col in TestResult:
|
|
52
|
+
if col not in counts.columns:
|
|
53
|
+
counts[col] = 0
|
|
54
|
+
counts = counts.groupby(["treatment", "outcome"]).sum().reset_index()[list(TestResult)]
|
|
55
|
+
# The below line normalises by the number of tests *for each edge*
|
|
56
|
+
# Independence tests X _||_ Y get two tests (X _||_ Y and Y _||_ X) because we don't know which way the
|
|
57
|
+
# causality flows. We need to normalise this (e.g. if X _||_ Y and Y _||_ X both pass, then the score should be
|
|
58
|
+
# 1 rather than 2) otherwise we end up unintentionally optimising for more independences.
|
|
59
|
+
counts = counts.apply(lambda col: col / counts.sum(axis=1))
|
|
60
|
+
|
|
61
|
+
return counts.sum(axis=0).to_dict()
|
|
62
|
+
|
|
63
|
+
def evaluate_fitness(
|
|
64
|
+
self,
|
|
65
|
+
individual: CausalDAG,
|
|
66
|
+
) -> tuple[tuple[float, float, float], list[tuple[str, str]]]:
|
|
67
|
+
"""
|
|
68
|
+
Evaluate the fitness of a given causal DAG by evaluating the corresponding test cases using a tier based
|
|
69
|
+
fitness metric.
|
|
70
|
+
lexicographical order (max pass, minimise failure, minimise unknown)
|
|
71
|
+
e.g. (X pass, Y fail, Z+1 unknown) is better than (X pass, Y+1 fail, Z unknown)
|
|
72
|
+
|
|
73
|
+
:param individual: The candidate individual to evaluate.
|
|
74
|
+
:returns: Tuple of the form (X, Y), where X is a triple containing the number of passing, failing, and
|
|
75
|
+
inestimable tests respectively, and Y is a list of failing edges.
|
|
76
|
+
"""
|
|
77
|
+
self.evaluate_tests(individual)
|
|
78
|
+
counts = self.sum_test_outcomes(individual.test_results)
|
|
79
|
+
|
|
80
|
+
# Add extra "var1" and "var2" columns to serve as order independent "treatment" and "outcome"
|
|
81
|
+
query_df = pd.concat(
|
|
82
|
+
[
|
|
83
|
+
individual.test_results,
|
|
84
|
+
pd.DataFrame(
|
|
85
|
+
np.sort(individual.test_results[["treatment", "outcome"]], axis=1), columns=["var1", "var2"]
|
|
86
|
+
),
|
|
87
|
+
],
|
|
88
|
+
axis=1,
|
|
89
|
+
)
|
|
90
|
+
problem_tests = query_df.groupby(["var1", "var2"]).filter(
|
|
91
|
+
# Groups are problematic if at least one test fails or no test passes
|
|
92
|
+
lambda group: (group["result"] == TestResult.FAIL).any()
|
|
93
|
+
or ~(group["result"] == TestResult.PASS).any()
|
|
94
|
+
)
|
|
95
|
+
problem_edges = problem_tests[["treatment", "outcome"]].apply(tuple, axis=1).tolist()
|
|
96
|
+
|
|
97
|
+
fitness_values = (
|
|
98
|
+
counts.get(TestResult.PASS, 0),
|
|
99
|
+
-counts.get(TestResult.FAIL, 0),
|
|
100
|
+
-counts.get(TestResult.INESTIMABLE, 0),
|
|
101
|
+
)
|
|
102
|
+
return fitness_values, problem_edges
|
|
103
|
+
|
|
104
|
+
def discover(self) -> CausalDAG:
|
|
105
|
+
"""
|
|
106
|
+
Discover the causal DAG.
|
|
107
|
+
|
|
108
|
+
:returns: The inferred causal DAG.
|
|
109
|
+
"""
|
|
110
|
+
|
|
111
|
+
start_time = time.time()
|
|
112
|
+
individual = CausalDAG()
|
|
113
|
+
individual.add_nodes_from(self.df.columns)
|
|
114
|
+
individual.add_edges_from(self.possible_edges)
|
|
115
|
+
self.remove_cycles(individual)
|
|
116
|
+
fitness_values, problem_edges = self.evaluate_fitness(individual)
|
|
117
|
+
|
|
118
|
+
iterations = self.max_iterations
|
|
119
|
+
iterations_without_improvement = 0
|
|
120
|
+
|
|
121
|
+
while problem_edges and iterations:
|
|
122
|
+
iterations -= 1
|
|
123
|
+
|
|
124
|
+
new_individual = individual.copy()
|
|
125
|
+
for origin, dest in random.sample(
|
|
126
|
+
# If we've gone over the maximum iterations without improvement
|
|
127
|
+
problem_edges
|
|
128
|
+
+ (
|
|
129
|
+
self.possible_edges
|
|
130
|
+
if iterations_without_improvement > self.max_iterations_without_improvement
|
|
131
|
+
else []
|
|
132
|
+
),
|
|
133
|
+
random.randint(1, len(problem_edges)),
|
|
134
|
+
):
|
|
135
|
+
if new_individual.has_edge(origin, dest) and (origin, dest) not in self.include_edges:
|
|
136
|
+
new_individual.remove_edge(origin, dest)
|
|
137
|
+
elif not new_individual.has_edge(origin, dest) and (origin, dest) not in self.exclude_edges:
|
|
138
|
+
# Want to bypass the cycle check of CausalDAG as we remove the cycles afterwards
|
|
139
|
+
new_individual.add_edge(origin, dest, ignore_cycles=True)
|
|
140
|
+
self.remove_cycles(new_individual)
|
|
141
|
+
new_fitness_values, new_problem_edges = self.evaluate_fitness(new_individual)
|
|
142
|
+
|
|
143
|
+
if new_fitness_values > fitness_values:
|
|
144
|
+
fitness_values = new_fitness_values
|
|
145
|
+
problem_edges = new_problem_edges
|
|
146
|
+
individual = new_individual
|
|
147
|
+
iterations_without_improvement = 0
|
|
148
|
+
else:
|
|
149
|
+
iterations_without_improvement += 1
|
|
150
|
+
|
|
151
|
+
end_time = time.time()
|
|
152
|
+
individual.graph["fitness"] = fitness_values
|
|
153
|
+
individual.graph["time"] = round(end_time - start_time)
|
|
154
|
+
|
|
155
|
+
return individual
|