In [1]:
from pathlib import Path
from survey_kit.utilities.random import RandomData
from survey_kit.utilities.formula_builder import FormulaBuilder
from survey_kit.calibration.moment import Moment
from survey_kit.calibration.calibration import Calibration
from survey_kit.utilities.dataframe import summary
import narwhals as nw
from survey_kit import logger
   WARNING:		pypardiso is unavailable (no usable MKL runtime on this CPU/platform); sparse solves will fall back to scikit-sparse/scipy, which are slower.
   WARNING:		sparse_dot_mkl is unavailable (no usable MKL runtime on this CPU/platform); falling back to scipy/numpy matrix products, which are slower.
In [2]:
logger.info("Generating data for weighting")
n_rows = 100_000
df_population = (
    RandomData(n_rows=n_rows, seed=12332151)
    .index("index")
    .integer("v_1", 1, 10)
    .np_distribution("v_f_continuous_0", "normal", loc=10, scale=2)
    .np_distribution("v_f_continuous_1", "normal", loc=10, scale=2)
    .np_distribution("v_f_continuous_2", "normal", loc=10, scale=2)
    .float("v_extra", -1, 2)
    .np_distribution("weight_0", "normal", loc=10, scale=1)
    .np_distribution("weight_1", "normal", loc=10, scale=1)
    .integer("year", 2016, 2021)
    .integer("month", 1, 12)
    .to_df()
    .lazy()
)

df_treatment = (
    RandomData(n_rows=n_rows, seed=894654)
    .index("index")
    .integer("v_1", 1, 10)
    #   Intentionally set the loc/scale as different than above
    .np_distribution("v_f_continuous_0", "normal", loc=11, scale=4)
    .np_distribution("v_f_continuous_1", "normal", loc=11, scale=4)
    .np_distribution("v_f_continuous_2", "normal", loc=11, scale=4)
    .float("v_extra", -1, 2)
    .np_distribution("weight_0", "normal", loc=10, scale=1)
    .np_distribution("weight_1", "normal", loc=10, scale=1)
    .integer("year", 2016, 2021)
    .integer("month", 1, 12)
    .to_df()
    .lazy()
)

# print(df.describe())
Generating data for weighting
In [3]:
logger.info("Weighting 'function'")
f = FormulaBuilder(df=df_population, constant=False)
f.continuous(columns=["v_1", "v_f_continuous_*", "v_f_p2_*"])
#   f.simple_interaction(columns=["v_1","v_f_continuous_0"])

logger.info("Define the target moments that the weighting will match")
logger.info("   This can be a dataset or a single row of pop controls")
m = Moment(
    df=df_population,
    formula=f.formula,
    weight="weight_0",
    index="index",
    by=["year"],
    equalize_by=True,
    rescale=True,
)

logger.info("You can save/reload moments if you want")
# m.save("/my/path/moment")
# m_loaded = Moment.load("/my/path/moment")
Weighting 'function'
Define the target moments that the weighting will match
   This can be a dataset or a single row of pop controls
You can save/reload moments if you want
In [4]:
#   Calibrate the data in df_treatment to the moment above
c = Calibration(
    df=df_treatment, moments=m, weight="weight_1", final_weight="weight_final"
)

c.run(
    #   Drop a moment if there are too few observations
    min_obs=5,
    # If it fails to converge, set bounds on the weights
    #   final weights = (base*ratio) where the bounds are on the ratio
    #   for "best possible" weights
    bounds=(0.001, 1000),
)

#   Merge the final weights back on the treatment data
df_treatment = c.get_final_weights(df_treatment)
Aggregating any sub_moments
Calibrating weights using aebw
      min obs = 5
     Calibration using combined moments
Entropy Balance Rewighting, Sanders (2024)
Input matrix is sparse? False
Problem Size: 100000 rows, 31 moments
  #    Criterion     ||Eq. Const.||  ||FOC Lagr.|| PrimalStepSize  DualStepSize  Opt. Violation.
SparseLinearSolver(size=31, spd=True) using: scipy
  0     0.000000      42831.275989       0.0000          inf            inf      42831.27598944
  1   2739.801893     18697.394213      56.3039     134.32575940    0.87719892   18697.47898708
  2   9503.564667      12.648764        10.3916      71.03346459    0.02828378    16.36999643  
  3   9483.511832       0.267324         1.2440      12.07772388    0.02872636     1.27236485  
  4   9483.161421       0.001644         0.3324      0.87395025     0.00072022     0.33245405  
  5   9483.150962       0.000008         0.0338      0.04439508     0.00000710     0.03383681  
  6   9483.150804       0.000000         0.0004      0.00531828     0.00000113     0.00039837  
  7   9483.150804       0.000000         0.0000      0.00006410     0.00000001     0.00000006  
Optimality converged.
Time elapsed: 0.3446362000031513
Optimization completed, success?: True

Converged:          True
Maximum Difference: 0.1134415812688645
shape: (30, 9)
┌────────────────────────────────┬──────────┬─────────┬───────────┬────────────┬─────────┬────────────────┬───────────┬───────────┐
│                       Variable ┆  Initial ┆ Targets ┆ Estimates ┆ Calibrated ┆ NonZero ┆ NonZero_Target ┆      diff ┆   percent │
╞════════════════════════════════╪══════════╪═════════╪═══════════╪════════════╪═════════╪════════════════╪═══════════╪═══════════╡
│ m0_year==2018:v_f_continuous_2 ┆ 1.113311 ┆       1 ┆  1.113442 ┆          1 ┆   16804 ┆          16550 ┆  0.113442 ┆ 11.344158 │
│ m0_year==2018:v_f_continuous_1 ┆  1.11314 ┆       1 ┆  1.111781 ┆          1 ┆   16804 ┆          16550 ┆  0.111781 ┆ 11.178072 │
│ m0_year==2018:v_f_continuous_0 ┆ 1.112906 ┆       1 ┆  1.108594 ┆          1 ┆   16804 ┆          16550 ┆  0.108594 ┆ 10.859389 │
│ m0_year==2021:v_f_continuous_0 ┆ 1.103274 ┆       1 ┆  1.103745 ┆          1 ┆   16681 ┆          16506 ┆  0.103745 ┆ 10.374487 │
│ m0_year==2019:v_f_continuous_1 ┆ 1.101078 ┆       1 ┆  1.102324 ┆          1 ┆   16624 ┆          16763 ┆  0.102324 ┆ 10.232387 │
│ m0_year==2017:v_f_continuous_2 ┆ 1.096367 ┆       1 ┆  1.102151 ┆          1 ┆   16569 ┆          16887 ┆  0.102151 ┆  10.21514 │
│ m0_year==2021:v_f_continuous_2 ┆  1.09976 ┆       1 ┆  1.101751 ┆          1 ┆   16681 ┆          16506 ┆  0.101751 ┆ 10.175121 │
│ m0_year==2019:v_f_continuous_0 ┆ 1.102635 ┆       1 ┆  1.101242 ┆          1 ┆   16624 ┆          16763 ┆  0.101242 ┆  10.12422 │
│ m0_year==2021:v_f_continuous_1 ┆  1.09959 ┆       1 ┆  1.100953 ┆          1 ┆   16681 ┆          16506 ┆  0.100953 ┆ 10.095282 │
│ m0_year==2017:v_f_continuous_0 ┆ 1.092225 ┆       1 ┆  1.100887 ┆          1 ┆   16569 ┆          16887 ┆  0.100887 ┆ 10.088701 │
│ m0_year==2017:v_f_continuous_1 ┆ 1.092783 ┆       1 ┆  1.100127 ┆          1 ┆   16569 ┆          16887 ┆  0.100127 ┆ 10.012654 │
│ m0_year==2020:v_f_continuous_2 ┆ 1.100242 ┆       1 ┆  1.098794 ┆          1 ┆   16646 ┆          16618 ┆  0.098794 ┆  9.879395 │
│ m0_year==2016:v_f_continuous_1 ┆  1.10321 ┆       1 ┆  1.098022 ┆          1 ┆   16676 ┆          16676 ┆  0.098022 ┆  9.802227 │
│ m0_year==2016:v_f_continuous_0 ┆ 1.101391 ┆       1 ┆  1.097971 ┆          1 ┆   16676 ┆          16676 ┆  0.097971 ┆  9.797117 │
│ m0_year==2020:v_f_continuous_1 ┆ 1.097323 ┆       1 ┆  1.094964 ┆          1 ┆   16646 ┆          16618 ┆  0.094964 ┆  9.496403 │
│ m0_year==2019:v_f_continuous_2 ┆  1.09556 ┆       1 ┆  1.094482 ┆          1 ┆   16624 ┆          16763 ┆  0.094482 ┆  9.448167 │
│ m0_year==2016:v_f_continuous_2 ┆  1.09744 ┆       1 ┆  1.093343 ┆          1 ┆   16676 ┆          16676 ┆  0.093343 ┆  9.334261 │
│ m0_year==2020:v_f_continuous_0 ┆ 1.092051 ┆       1 ┆  1.092315 ┆          1 ┆   16646 ┆          16618 ┆  0.092315 ┆  9.231494 │
│              m0_year==2018:v_1 ┆  1.01166 ┆       1 ┆  1.013113 ┆          1 ┆   16804 ┆          16550 ┆  0.013113 ┆  1.311349 │
│              m0_year==2021:v_1 ┆ 1.006379 ┆       1 ┆  1.009274 ┆          1 ┆   16681 ┆          16506 ┆  0.009274 ┆  0.927371 │
│              m0_year==2018:_in ┆ 1.008746 ┆       1 ┆  1.007851 ┆          1 ┆   16804 ┆          16550 ┆  0.007851 ┆  0.785068 │
│              m0_year==2016:v_1 ┆ 0.999911 ┆       1 ┆  0.992515 ┆          1 ┆   16676 ┆          16676 ┆ -0.007485 ┆ -0.748549 │
│              m0_year==2019:v_1 ┆ 1.008309 ┆       1 ┆   1.00695 ┆          1 ┆   16624 ┆          16763 ┆   0.00695 ┆  0.694964 │
│              m0_year==2017:v_1 ┆ 1.002875 ┆       1 ┆    1.0059 ┆          1 ┆   16569 ┆          16887 ┆    0.0059 ┆  0.590018 │
│              m0_year==2019:_in ┆ 0.997436 ┆       1 ┆  0.995725 ┆          1 ┆   16624 ┆          16763 ┆ -0.004275 ┆ -0.427468 │
│              m0_year==2016:_in ┆ 1.000639 ┆       1 ┆  0.996354 ┆          1 ┆   16676 ┆          16676 ┆ -0.003646 ┆ -0.364649 │
│              m0_year==2021:_in ┆ 1.001684 ┆       1 ┆  1.002978 ┆          1 ┆   16681 ┆          16506 ┆  0.002978 ┆  0.297803 │
│              m0_year==2020:_in ┆ 0.998146 ┆       1 ┆  0.997146 ┆          1 ┆   16646 ┆          16618 ┆ -0.002854 ┆ -0.285395 │
│              m0_year==2020:v_1 ┆ 1.001671 ┆       1 ┆  1.001688 ┆          1 ┆   16646 ┆          16618 ┆  0.001688 ┆  0.168779 │
│              m0_year==2017:_in ┆ 0.993348 ┆       1 ┆  0.999946 ┆          1 ┆   16569 ┆          16887 ┆ -0.000054 ┆ -0.005359 │
└────────────────────────────────┴──────────┴─────────┴───────────┴────────────┴─────────┴────────────────┴───────────┴───────────┘
In [5]:
logger.info("'Population' estimates")
_ = summary(df_population, weight="weight_0")
'Population' estimates
┌──────────────────┬─────────┬─────────────┬───────────────┬───────────────┬───────────┬───────────┐
│         Variable ┆       n ┆ n (missing) ┆          mean ┆           std ┆       min ┆       max │
╞══════════════════╪═════════╪═════════════╪═══════════════╪═══════════════╪═══════════╪═══════════╡
│            index ┆ 100,000 ┆           0 ┆ 49,987.162827 ┆ 28,870.556017 ┆       0.0 ┆  99,999.0 │
│              v_1 ┆ 100,000 ┆           0 ┆      5.484124 ┆      2.875588 ┆       1.0 ┆      10.0 │
│ v_f_continuous_0 ┆ 100,000 ┆           0 ┆      9.989362 ┆      1.996633 ┆  1.491748 ┆ 18.835062 │
│ v_f_continuous_1 ┆ 100,000 ┆           0 ┆     10.002072 ┆      2.006889 ┆   1.37638 ┆  18.82769 │
│ v_f_continuous_2 ┆ 100,000 ┆           0 ┆      9.998039 ┆      2.004505 ┆  1.166252 ┆ 19.254231 │
│          v_extra ┆ 100,000 ┆           0 ┆      0.504297 ┆      0.867154 ┆ -0.999978 ┆  1.999996 │
│         weight_1 ┆ 100,000 ┆           0 ┆     10.004412 ┆      1.002786 ┆   5.35507 ┆ 14.023032 │
│             year ┆ 100,000 ┆           0 ┆  2,018.493399 ┆      1.706467 ┆   2,016.0 ┆   2,021.0 │
│            month ┆ 100,000 ┆           0 ┆      6.504686 ┆       3.44999 ┆       1.0 ┆      12.0 │
└──────────────────┴─────────┴─────────────┴───────────────┴───────────────┴───────────┴───────────┘
In [6]:
logger.info("\n\n'Treatment', original weights")
_ = summary(df_treatment, weight="weight_1")

'Treatment', original weights
┌──────────────────┬─────────┬─────────────┬───────────────┬───────────────┬───────────┬───────────┐
│         Variable ┆       n ┆ n (missing) ┆          mean ┆           std ┆       min ┆       max │
╞══════════════════╪═════════╪═════════════╪═══════════════╪═══════════════╪═══════════╪═══════════╡
│            index ┆ 100,000 ┆           0 ┆ 49,997.423141 ┆ 28,873.734839 ┆       0.0 ┆  99,999.0 │
│              v_1 ┆ 100,000 ┆           0 ┆       5.51228 ┆      2.868293 ┆       1.0 ┆      10.0 │
│ v_f_continuous_0 ┆ 100,000 ┆           0 ┆      10.99576 ┆      4.013016 ┆ -5.459849 ┆ 30.287625 │
│ v_f_continuous_1 ┆ 100,000 ┆           0 ┆     11.014155 ┆      3.999313 ┆ -7.360353 ┆ 27.415263 │
│ v_f_continuous_2 ┆ 100,000 ┆           0 ┆      11.00231 ┆      4.005346 ┆ -8.478953 ┆ 27.102181 │
│          v_extra ┆ 100,000 ┆           0 ┆      0.500786 ┆      0.867307 ┆ -0.999996 ┆  1.999988 │
│         weight_0 ┆ 100,000 ┆           0 ┆      9.997927 ┆      1.001265 ┆   5.26969 ┆ 14.090873 │
│             year ┆ 100,000 ┆           0 ┆  2,018.500692 ┆      1.707684 ┆   2,016.0 ┆   2,021.0 │
│            month ┆ 100,000 ┆           0 ┆      6.489706 ┆      3.447587 ┆       1.0 ┆      12.0 │
│     weight_final ┆ 100,000 ┆           0 ┆      1.010076 ┆      0.474715 ┆  0.126062 ┆  7.945797 │
└──────────────────┴─────────┴─────────────┴───────────────┴───────────────┴───────────┴───────────┘
In [7]:
logger.info("\n\n'Treatment', calibrated")
_ = summary(df_treatment, weight="weight_final")

'Treatment', calibrated
┌──────────────────┬─────────┬─────────────┬───────────────┬───────────────┬───────────┬───────────┐
│         Variable ┆       n ┆ n (missing) ┆          mean ┆           std ┆       min ┆       max │
╞══════════════════╪═════════╪═════════════╪═══════════════╪═══════════════╪═══════════╪═══════════╡
│            index ┆ 100,000 ┆           0 ┆ 49,987.572128 ┆ 28,868.167031 ┆       0.0 ┆  99,999.0 │
│              v_1 ┆ 100,000 ┆           0 ┆      5.511032 ┆       2.86937 ┆       1.0 ┆      10.0 │
│ v_f_continuous_0 ┆ 100,000 ┆           0 ┆     10.996213 ┆      4.014495 ┆ -5.459849 ┆ 30.287625 │
│ v_f_continuous_1 ┆ 100,000 ┆           0 ┆     11.015899 ┆      4.000952 ┆ -7.360353 ┆ 27.415263 │
│ v_f_continuous_2 ┆ 100,000 ┆           0 ┆     11.004445 ┆      4.003024 ┆ -8.478953 ┆ 27.102181 │
│          v_extra ┆ 100,000 ┆           0 ┆      0.500065 ┆      0.867396 ┆ -0.999996 ┆  1.999988 │
│         weight_0 ┆ 100,000 ┆           0 ┆     10.000925 ┆      1.002408 ┆   5.26969 ┆ 14.090873 │
│         weight_1 ┆ 100,000 ┆           0 ┆      10.10284 ┆      0.999476 ┆  5.539069 ┆ 14.227871 │
│             year ┆ 100,000 ┆           0 ┆   2,018.50105 ┆      1.707354 ┆   2,016.0 ┆   2,021.0 │
│            month ┆ 100,000 ┆           0 ┆      6.493253 ┆      3.446045 ┆       1.0 ┆      12.0 │
└──────────────────┴─────────┴─────────────┴───────────────┴───────────────┴───────────┴───────────┘