{
"cells": [
{
"cell_type": "markdown",
"id": "7277b9c2",
"metadata": {},
"source": [
"# Genetic Matcher\n",
"\n",
"The GeneticMatcher can be used to optimize any function of the baseline covariates, both linear and non-linear. In this demo notebook, we show how to call the matcher in the PyBalance library, including an example of a non-linear balance function."
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "5d3c98f7",
"metadata": {},
"outputs": [],
"source": [
"import logging \n",
"logging.basicConfig(\n",
" format=\"%(levelname)-4s [%(filename)s:%(lineno)d] %(message)s\",\n",
" level='INFO',\n",
")\n",
"\n",
"from pybalance.sim import generate_toy_dataset\n",
"from pybalance.utils import (\n",
" BetaBalance, \n",
" BetaSquaredBalance, \n",
" BetaXBalance,\n",
" BetaMaxBalance,\n",
" GammaBalance, \n",
" GammaSquaredBalance,\n",
" GammaXBalance,\n",
" GammaXTreeBalance,\n",
" MatchingData\n",
")\n",
"from pybalance.genetic import GeneticMatcher, get_global_defaults\n",
"from pybalance.visualization import (\n",
" plot_numeric_features, \n",
" plot_categoric_features, \n",
" plot_binary_features,\n",
" plot_per_feature_loss,\n",
")\n",
"\n",
"time_limit = 120"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "0df76346",
"metadata": {},
"outputs": [
{
"data": {
"text/html": [
"\n",
" Headers Numeric:
\n",
" ['age', 'height', 'weight']
\n",
" Headers Categoric:
\n",
" ['gender', 'haircolor', 'country', 'binary_0', 'binary_1', 'binary_2', 'binary_3']
\n",
" Populations
\n",
" ['pool', 'target']
\n",
"
| \n", " | age | \n", "height | \n", "weight | \n", "gender | \n", "haircolor | \n", "country | \n", "population | \n", "binary_0 | \n", "binary_1 | \n", "binary_2 | \n", "binary_3 | \n", "patient_id | \n", "
|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | \n", "62.511573 | \n", "190.229250 | \n", "105.165097 | \n", "0.0 | \n", "2 | \n", "3 | \n", "pool | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "
| 1 | \n", "68.505065 | \n", "161.121236 | \n", "95.001474 | \n", "0.0 | \n", "1 | \n", "1 | \n", "pool | \n", "1 | \n", "0 | \n", "1 | \n", "0 | \n", "1 | \n", "
| 2 | \n", "50.071384 | \n", "162.325356 | \n", "84.290576 | \n", "1.0 | \n", "0 | \n", "5 | \n", "pool | \n", "0 | \n", "0 | \n", "1 | \n", "1 | \n", "2 | \n", "
| 3 | \n", "44.423692 | \n", "150.948096 | \n", "82.031381 | \n", "1.0 | \n", "2 | \n", "2 | \n", "pool | \n", "0 | \n", "0 | \n", "0 | \n", "1 | \n", "3 | \n", "
| 4 | \n", "41.695052 | \n", "132.952651 | \n", "54.857540 | \n", "0.0 | \n", "1 | \n", "3 | \n", "pool | \n", "0 | \n", "0 | \n", "1 | \n", "1 | \n", "4 | \n", "
| ... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "
| 995 | \n", "21.474205 | \n", "168.602546 | \n", "70.342128 | \n", "0.0 | \n", "2 | \n", "5 | \n", "target | \n", "0 | \n", "0 | \n", "0 | \n", "1 | \n", "10995 | \n", "
| 996 | \n", "40.643320 | \n", "188.188724 | \n", "61.611744 | \n", "0.0 | \n", "2 | \n", "4 | \n", "target | \n", "1 | \n", "0 | \n", "0 | \n", "1 | \n", "10996 | \n", "
| 997 | \n", "29.472765 | \n", "161.408162 | \n", "57.214095 | \n", "0.0 | \n", "0 | \n", "1 | \n", "target | \n", "0 | \n", "1 | \n", "1 | \n", "1 | \n", "10997 | \n", "
| 998 | \n", "41.291949 | \n", "150.968833 | \n", "91.270798 | \n", "0.0 | \n", "0 | \n", "3 | \n", "target | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "10998 | \n", "
| 999 | \n", "67.530294 | \n", "155.124741 | \n", "56.196505 | \n", "1.0 | \n", "0 | \n", "1 | \n", "target | \n", "1 | \n", "0 | \n", "0 | \n", "0 | \n", "10999 | \n", "
11000 rows × 12 columns
\n", "| \n", " | age | \n", "height | \n", "weight | \n", "gender | \n", "haircolor | \n", "country | \n", "population | \n", "binary_0 | \n", "binary_1 | \n", "binary_2 | \n", "binary_3 | \n", "patient_id | \n", "
|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 0 | \n", "55.261578 | \n", "139.396134 | \n", "94.438359 | \n", "0.0 | \n", "2 | \n", "2 | \n", "target | \n", "0 | \n", "0 | \n", "1 | \n", "1 | \n", "10000 | \n", "
| 1 | \n", "63.113091 | \n", "165.563337 | \n", "67.433016 | \n", "1.0 | \n", "2 | \n", "2 | \n", "target | \n", "0 | \n", "1 | \n", "1 | \n", "0 | \n", "10001 | \n", "
| 2 | \n", "58.232216 | \n", "160.859857 | \n", "71.915385 | \n", "1.0 | \n", "0 | \n", "2 | \n", "target | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "10002 | \n", "
| 3 | \n", "58.996941 | \n", "140.357415 | \n", "115.606615 | \n", "1.0 | \n", "0 | \n", "3 | \n", "target | \n", "1 | \n", "1 | \n", "0 | \n", "0 | \n", "10003 | \n", "
| 4 | \n", "36.850195 | \n", "189.983706 | \n", "53.000581 | \n", "0.0 | \n", "2 | \n", "5 | \n", "target | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "10004 | \n", "
| ... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "... | \n", "
| 5044 | \n", "42.548928 | \n", "129.729442 | \n", "94.445375 | \n", "1.0 | \n", "1 | \n", "2 | \n", "pool | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "5044 | \n", "
| 1144 | \n", "29.400226 | \n", "167.737236 | \n", "76.118095 | \n", "1.0 | \n", "0 | \n", "4 | \n", "pool | \n", "0 | \n", "1 | \n", "0 | \n", "1 | \n", "1144 | \n", "
| 5314 | \n", "50.104985 | \n", "163.663484 | \n", "85.785445 | \n", "1.0 | \n", "2 | \n", "4 | \n", "pool | \n", "0 | \n", "1 | \n", "0 | \n", "0 | \n", "5314 | \n", "
| 2174 | \n", "54.372402 | \n", "149.801277 | \n", "92.946485 | \n", "0.0 | \n", "1 | \n", "2 | \n", "pool | \n", "0 | \n", "0 | \n", "0 | \n", "0 | \n", "2174 | \n", "
| 8610 | \n", "72.912497 | \n", "185.908237 | \n", "96.553338 | \n", "0.0 | \n", "1 | \n", "4 | \n", "pool | \n", "1 | \n", "1 | \n", "1 | \n", "1 | \n", "8610 | \n", "
2000 rows × 12 columns
\n", "