In [1]:
import numpy as np
import polars as pl
from survey_kit.imputation.variable import Variable
from survey_kit.imputation.srmi import SRMI
from survey_kit import logger, config
In [2]:
# Draw data with a categorical PREDICTOR (not itself being imputed) and a
# grouping variable (state) whose effect we want captured without treating
# it as a plain dummy-coded predictor
n_rows = 4_000
rng = np.random.default_rng(20260913)
x1 = rng.normal(size=n_rows)
industry_levels = ["retail", "healthcare", "manufacturing", "tech"]
industry_effect = {"retail": -2_000.0, "healthcare": 3_000.0, "manufacturing": 0.0, "tech": 6_000.0}
industry = rng.choice(industry_levels, size=n_rows)
state_levels = ["ca", "tx", "ny", "fl"]
state_effect = {"ca": 4_000.0, "tx": -1_000.0, "ny": 2_000.0, "fl": -2_000.0}
state = rng.choice(state_levels, size=n_rows)
income = (
40_000
+ 8_000 * x1
+ np.array([industry_effect[i] for i in industry])
+ np.array([state_effect[s] for s in state])
+ rng.normal(scale=4_000, size=n_rows)
)
df = pl.DataFrame(
dict(
person_id=range(n_rows),
x1=x1,
industry=industry,
state=state,
income=income,
)
)
missing_income = rng.random(n_rows) < 0.2
df = df.with_columns(
pl.when(pl.Series(missing_income)).then(None).otherwise(pl.col("income")).alias("income")
)
In [3]:
logger.info(
"'industry' is a string predictor for income, not something being imputed - "
"declare it via categorical_predictors so the model treats it as a real "
"category, not a scrambled numeric column"
)
logger.info(
"'state' is a grouping variable - group_levels lets the model borrow strength "
"across observations that share a state, without adding 'state' as an ordinary "
"dummy-coded predictor"
)
srmi = SRMI.simple_model(
df=df,
index="person_id",
categorical_predictors=["industry"],
group_levels="state",
# group_levels only actually applies to models that support it -
# the continuous default (LightGBM) doesn't, so switch to
# RandomForest here to see it take effect
model={Variable.Class.continuous: Variable.ModelType.RandomForest},
replication=SRMI.Replication(n_implicates=2, n_iterations=1),
parallel=SRMI.Parallel(enabled=False),
storage=SRMI.Storage(
path_model=f"{config.path_temp_files}/tutorial_simple_model_categorical_group",
force_start=True,
),
)
v_income = srmi.variables[0]
logger.info(f"income predictors: {v_income.model}")
logger.info(f"income categorical_feature: {v_income.parameters.get('categorical_feature')}")
logger.info(f"income group_levels: {v_income.parameters.get('group_levels')}")
'industry' is a string predictor for income, not something being imputed - declare it via categorical_predictors so the model treats it as a real category, not a scrambled numeric column
'state' is a grouping variable - group_levels lets the model borrow strength across observations that share a state, without adding 'state' as an ordinary dummy-coded predictor
auto_detect: 'income' -> class=continuous, modeltype=RandomForest, predictors=['x1', 'industry']
Removing existing directory C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.scratch\temp_files/tutorial_simple_model_categorical_group.srmi
income predictors: ~1+x1+C(industry)
income categorical_feature: []
income group_levels: ['state']
In [4]:
logger.info("Run it")
srmi.run()
Run it
Variable selection before SRMI run, if necessary
income: Method.No
Hyperparameter tuning before SRMI run, if necessary
Removing existing directory C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.scratch\temp_files/tutorial_simple_model_categorical_group.srmi/1.srmi.implicate
Removing existing directory C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.scratch\temp_files/tutorial_simple_model_categorical_group.srmi/2.srmi.implicate
Imputation using RandomForest
C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.venv\Lib\site-packages\sklearn\base.py:1365: DataConversionWarning: A column-vector y was passed when a 1d array was expected. Please change the shape of y to (n_samples,), for example using ravel(). return fit_method(estimator, *args, **kwargs)
R2 = 0.9664
┌──────────────────────────┬──────────┐ │ Variable ┆ Beta │ ╞══════════════════════════╪══════════╡ │ x1 ┆ 0.8917 │ │ C(industry)healthcare ┆ 0.01941 │ │ C(industry)manufacturing ┆ 0.006627 │ │ C(industry)retail ┆ 0.01841 │ │ C(industry)tech ┆ 0.06381 │ └──────────────────────────┴──────────┘
error=pmm: donating observed value(s) ['income'] from 10-nearest matched donors
Finding 10 nearest neighbors on ['___prediction']
Randomly picking one and donating ['income']
Most common matches:
shape: (5, 2) ┌───────────┬─────────┐ │ person_id ┆ nDonors │ │ --- ┆ --- │ │ i16 ┆ i8 │ ╞═══════════╪═════════╡ │ 2037 ┆ 5 │ │ 520 ┆ 4 │ │ 111 ┆ 3 │ │ 128 ┆ 3 │ │ 1136 ┆ 3 │ └───────────┴─────────┘
Post-imputation statistics for ['income']
Where: None
Where (impute): col(___imp_missing_income_1)
┌──────────┬─────────┬──────┬──────────────┬─────────┬────────┬──────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┐ │ Variable ┆ Imputed ┆ n ┆ n (not null) ┆ mean ┆ std ┆ mean (not 0) ┆ std (not 0) ┆ q10 (not 0) ┆ q25 (not 0) ┆ q50 (not 0) ┆ q75 (not 0) ┆ q90 (not 0) ┆ min (not 0) ┆ max (not 0) │ ╞══════════╪═════════╪══════╪══════════════╪═════════╪════════╪══════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╡ │ income ┆ ┆ 4000 ┆ 4000 ┆ 42750.0 ┆ 9881.0 ┆ 42750.0 ┆ 9881.0 ┆ 30020.0 ┆ 36100.0 ┆ 42810.0 ┆ 49050.0 ┆ 55660.0 ┆ 4395.0 ┆ 74220.0 │ │ income ┆ 0 ┆ 3180 ┆ 3180 ┆ 42670.0 ┆ 9892.0 ┆ 42670.0 ┆ 9892.0 ┆ 29960.0 ┆ 36110.0 ┆ 42720.0 ┆ 49030.0 ┆ 55430.0 ┆ 4395.0 ┆ 74220.0 │ │ income ┆ 1 ┆ 820 ┆ 820 ┆ 43040.0 ┆ 9840.0 ┆ 43040.0 ┆ 9840.0 ┆ 30230.0 ┆ 35860.0 ┆ 43050.0 ┆ 49280.0 ┆ 56060.0 ┆ 12530.0 ┆ 71520.0 │ └──────────┴─────────┴──────┴──────────────┴─────────┴────────┴──────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┘
income
Final Estimates by Iteration
Removing existing directory C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.scratch\temp_files/tutorial_simple_model_categorical_group.srmi/1.srmi.implicate
Imputation using RandomForest
C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.venv\Lib\site-packages\sklearn\base.py:1365: DataConversionWarning: A column-vector y was passed when a 1d array was expected. Please change the shape of y to (n_samples,), for example using ravel(). return fit_method(estimator, *args, **kwargs)
R2 = 0.9659
┌──────────────────────────┬──────────┐ │ Variable ┆ Beta │ ╞══════════════════════════╪══════════╡ │ x1 ┆ 0.8908 │ │ C(industry)healthcare ┆ 0.02323 │ │ C(industry)manufacturing ┆ 0.006389 │ │ C(industry)retail ┆ 0.01308 │ │ C(industry)tech ┆ 0.06647 │ └──────────────────────────┴──────────┘
error=pmm: donating observed value(s) ['income'] from 10-nearest matched donors
Finding 10 nearest neighbors on ['___prediction']
Randomly picking one and donating ['income']
Most common matches:
shape: (5, 2) ┌───────────┬─────────┐ │ person_id ┆ nDonors │ │ --- ┆ --- │ │ i16 ┆ i8 │ ╞═══════════╪═════════╡ │ 755 ┆ 4 │ │ 2264 ┆ 4 │ │ 4 ┆ 3 │ │ 493 ┆ 3 │ │ 969 ┆ 3 │ └───────────┴─────────┘
Post-imputation statistics for ['income']
Where: None
Where (impute): col(___imp_missing_income_1)
┌──────────┬─────────┬──────┬──────────────┬─────────┬────────┬──────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┬─────────────┐ │ Variable ┆ Imputed ┆ n ┆ n (not null) ┆ mean ┆ std ┆ mean (not 0) ┆ std (not 0) ┆ q10 (not 0) ┆ q25 (not 0) ┆ q50 (not 0) ┆ q75 (not 0) ┆ q90 (not 0) ┆ min (not 0) ┆ max (not 0) │ ╞══════════╪═════════╪══════╪══════════════╪═════════╪════════╪══════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╪═════════════╡ │ income ┆ ┆ 4000 ┆ 4000 ┆ 42790.0 ┆ 9853.0 ┆ 42790.0 ┆ 9853.0 ┆ 30110.0 ┆ 36270.0 ┆ 42830.0 ┆ 49080.0 ┆ 55550.0 ┆ 4395.0 ┆ 74220.0 │ │ income ┆ 0 ┆ 3180 ┆ 3180 ┆ 42670.0 ┆ 9892.0 ┆ 42670.0 ┆ 9892.0 ┆ 29960.0 ┆ 36110.0 ┆ 42720.0 ┆ 49030.0 ┆ 55430.0 ┆ 4395.0 ┆ 74220.0 │ │ income ┆ 1 ┆ 820 ┆ 820 ┆ 43220.0 ┆ 9692.0 ┆ 43220.0 ┆ 9692.0 ┆ 30710.0 ┆ 36600.0 ┆ 43440.0 ┆ 49470.0 ┆ 56040.0 ┆ 13160.0 ┆ 71120.0 │ └──────────┴─────────┴──────┴──────────────┴─────────┴────────┴──────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┴─────────────┘
income
Final Estimates by Iteration
Removing existing directory C:\Users\jonro\OneDrive\Documents\Coding\survey_kit\.scratch\temp_files/tutorial_simple_model_categorical_group.srmi/2.srmi.implicate