Source code for pybalance.weighting.matcher

from __future__ import annotations

from typing import Dict, List, Optional, Tuple

import matplotlib
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import seaborn as sns
from scipy.optimize import minimize
from sklearn.base import BaseEstimator, clone
from sklearn.linear_model import LogisticRegression
from sklearn.preprocessing import StandardScaler

from pybalance.utils import MatchingData, BalanceCalculator

import logging

logger = logging.getLogger(__name__)


def _check_fitted(weighter):
    if weighter.weights is None:
        raise ValueError("Weighter has not been fitted!")


[docs]def effective_sample_size(weights: np.ndarray) -> float: """ Kish's effective sample size: ``(sum w)**2 / sum(w**2)``. Invariant to the overall scale of the weights. A large drop relative to ``len(weights)`` indicates the reweighting is relying heavily on a small number of pool patients to match the target, and results should be interpreted cautiously (e.g. the target may lie outside the range of the pool's covariates). """ weights = np.asarray(weights, dtype=float) return float(weights.sum() ** 2 / np.sum(weights**2))
def _softmax_weights(Z: np.ndarray, lam: np.ndarray) -> np.ndarray: linpred = Z @ lam linpred = linpred - linpred.max() w = np.exp(linpred) return w / w.sum() def _solve_entropy_weights( Z: np.ndarray, max_iter: int = 200, tol: float = 1e-10, ridge: float = 1e-8, penalty: Optional[np.ndarray] = None, ) -> Tuple[np.ndarray, Dict]: """ Solve for weights ``w_i = softmax(Z @ lambda)_i`` (so ``sum_i w_i == 1`` by construction) by minimizing the (convex, smooth) dual of the maximum-entropy weighting problem, ``log(sum_i exp(Z_i @ lambda))``, via scipy's trust-region-Newton-CG solver using an analytic gradient and Hessian. At the optimum, ``sum_i w_i * Z_i == 0`` for every column of ``Z``, i.e. every (centered) moment constraint is exactly balanced in the weighted pool. This is the method-of-moments weighting scheme underlying MAIC (Signorovitch et al., 2010) and entropy balancing (Hainmueller, 2012). The softmax parametrization (rather than the more commonly-quoted raw ``w_i = exp(z_i . lambda)``) is important for numerical robustness: the raw form leaves an unconstrained degree of freedom along which ``sum(w)`` can be driven towards 0 (or diverge) while the *raw* gradient ``Z.T @ w`` shrinks in lockstep, which can fool a naive convergence check into declaring success at a degenerate, non-solution point. Because the softmax always sums to 1, this failure mode is eliminated and the gradient (a genuine weighted mean) is a directly interpretable, scale-free convergence criterion. A trust-region solver (rather than plain damped Newton) is used because the Hessian -- a weighted covariance matrix of ``Z`` -- can be near-singular along directions with little effective curvature (e.g. highly imbalanced categoric rates), which otherwise produces enormous, unstable Newton steps. :param Z: (n_pool, n_constraints) matrix of centered, column-scaled moment constraints: ``Z[i, j]`` is patient ``i``'s (scaled) deviation from the target moment for constraint ``j``. :param max_iter: Maximum number of solver iterations. :param tol: Convergence tolerance on the max-norm of the constraint violation (weighted-mean residual). :param ridge: Ridge term added to the Hessian for numerical stability; important when constraints are collinear or near-collinear. :param penalty: Optional per-constraint soft-constraint strength. A positive entry adds ``penalty_j * lambda_j**2 / 2`` to the dual, which relaxes constraint ``j`` to a quadratic penalty on its residual instead of an exact equality. Needed for constraints sitting on the boundary of what positive weights can reach (e.g. "no patient above a max"), for which exact balance has no finite solution. :return: Tuple ``(weights, diagnostics)`` where diagnostics has keys 'converged', 'n_iter', 'max_constraint_violation' (over the exact constraints) and, if ``penalty`` is used, 'max_soft_violation'. """ n, k = Z.shape if k == 0: return np.ones(n) / n, { "converged": True, "n_iter": 0, "max_constraint_violation": 0.0, } penalty = np.zeros(k) if penalty is None else np.asarray(penalty, dtype=float) def _objective(lam): linpred = Z @ lam m = linpred.max() return m + np.log(np.exp(linpred - m).sum()) + 0.5 * np.sum(penalty * lam**2) def _grad(lam): return Z.T @ _softmax_weights(Z, lam) + penalty * lam def _hess(lam): w = _softmax_weights(Z, lam) zc = Z - (Z.T @ w) return (zc * w[:, None]).T @ zc + np.diag(penalty) + ridge * np.eye(k) res = minimize( _objective, x0=np.zeros(k), jac=_grad, hess=_hess, method="trust-ncg", options={"maxiter": max_iter, "gtol": tol}, ) w_final = _softmax_weights(Z, res.x) residual = np.abs(Z.T @ w_final) exact = penalty == 0 max_violation = float(residual[exact].max()) if exact.any() else 0.0 converged = bool(res.success) or max_violation < tol * 100 diagnostics = { "converged": converged, "n_iter": int(res.nit), "max_constraint_violation": max_violation, } if (~exact).any(): diagnostics["max_soft_violation"] = float(residual[~exact].max()) return w_final, diagnostics
[docs]class BaseWeighter: """ Common interface for weighting methods. Unlike the Matcher classes (genetic/lp/propensity), a Weighter never drops pool patients; instead it assigns every pool patient a non-negative weight so that the *weighted* pool resembles the target on the moments of interest. This is the standard approach for indirect/external comparisons when excluding patients is undesirable or infeasible -- e.g. small pools, or a target known only through published aggregate statistics. :param matching_data: MatchingData whose pool is to be weighted. The target can be either patient-level or an ``AggregateTarget`` (e.g. a published Table 1). :param weight_col: Name of the column used to store weights on the MatchingData returned by match(). Must not collide with an existing matching feature. :param verbose: Whether to log fitting diagnostics. """ def __init__( self, matching_data: MatchingData, weight_col: str = "sample_weight", verbose: bool = True, ): if weight_col in matching_data.headers.all: raise ValueError( f"weight_col={weight_col!r} collides with an existing matching " "feature. Pass a different weight_col." ) self.matching_data = matching_data.copy() self.weight_col = weight_col self.verbose = verbose self.weights: Optional[np.ndarray] = None self.diagnostics: Dict = {}
[docs] def get_params(self) -> Dict: """Return the weighter's configuration parameters as a dict.""" raise NotImplementedError
def _fit(self) -> "BaseWeighter": """ Compute self.weights/self.diagnostics. Subclasses implement this; end users should call match() instead (see below), which calls this internally and wraps the result in a MatchingData, matching the public interface of the other matchers in pybalance. """ raise NotImplementedError
[docs] def get_weights(self) -> np.ndarray: """ Return the fitted per-patient pool weights (in pool row order). """ _check_fitted(self) return self.weights
[docs] def effective_sample_size(self) -> float: """ Kish's effective sample size of the fitted weights. See ``effective_sample_size()``. """ _check_fitted(self) return effective_sample_size(self.weights)
[docs] def match(self) -> MatchingData: """ Fit (if not already fit) and return a MatchingData instance in which every pool patient is retained and carries a new column (see weight_col) holding their fitted weight. The target population (patient-level or aggregate) is passed through unchanged, with weight 1.0 assigned to any patient-level target rows. """ if self.weights is None: self._fit() md = self.matching_data pool = md.get_population(md.pool_name).copy() pool[self.weight_col] = self.weights if md.has_aggregate_target: return MatchingData( pool=pool, target=md.aggregate_target, headers=md.headers, population_col=md.population_col, pool_name=md.pool_name, target_name=md.target_name, ) target = md.get_population(md.target_name).copy() target[self.weight_col] = 1.0 return MatchingData( pool=pool, target=target, headers=md.headers, population_col=md.population_col, pool_name=md.pool_name, target_name=md.target_name, )
[docs]class EntropyBalanceWeighter(BaseWeighter): """ General maximum-entropy ("method of moments") weighting: solves for pool weights of the form ``w_i = exp(z_i . alpha)`` such that the weighted pool matches the target. For an ``AggregateTarget`` every disclosed statistic is its own constraint and nothing else is constrained: a numeric mean, the rate of each disclosed categoric level, and the rate above each disclosed median / quantile (the feature is dichotomized at the disclosed value). A feature or categoric level the target says nothing about is simply left free. Optionally also balances variance for numeric features whose target discloses (or, for a patient-level target, has) a standard deviation; for an ``AggregateTarget`` that needs the mean too, since the variance is taken around it. A disclosed ``min`` / ``max`` is a *soft* constraint, as in ``AggregateConstraintSatisfactionMatcher``: exactly zero weight on patients above a max cannot be reached by positive weights, so the fraction of weight there is only penalized (see ``limit_penalty``). This is the general form of Matching-Adjusted Indirect Comparison (MAIC; Signorovitch et al., 2010); see ``MAICWeighter`` for the classic mean-only formulation. :param matching_data: MatchingData whose pool is to be weighted. :param match_variance: If True, additionally constrain the weighted variance of numeric features. For an ``AggregateTarget``, only features that actually disclose a "std" are constrained (mirroring ``AggregateConstraintSatisfactionMatcher``); for a patient-level target, every numeric feature's variance is constrained. Categoric features are never variance-constrained: their variance is already determined by their (matched) rate. :param normalize: How to rescale the fitted weights for reporting: "target" (default) rescales so weights sum to the target population size, "pool" rescales so weights sum to the pool size, "none" leaves the raw dual-optimizer weights unscaled. This choice has no effect on balance (the moment constraints are scale-invariant) or on effective_sample_size(); it only affects the units the weights are reported in. :param max_iter: Maximum number of Newton iterations. :param tol: Convergence tolerance on the (max-norm) constraint violation. :param ridge: Ridge regularization added to the Newton step for numerical stability; increase this if fitting fails to converge due to collinear/near-collinear features. :param weight_col: Name of the column used to store weights on the MatchingData returned by match(). :param verbose: Whether to log fitting diagnostics. :param limit_penalty: Strength of the quadratic penalty on the weight placed beyond a disclosed min / max of an ``AggregateTarget`` (a soft constraint). Larger values enforce the limit more tightly. """ def __init__( self, matching_data: MatchingData, match_variance: bool = False, normalize: str = "target", max_iter: int = 200, tol: float = 1e-8, ridge: float = 1e-8, weight_col: str = "sample_weight", verbose: bool = True, limit_penalty: float = 1e-2, ): super().__init__(matching_data, weight_col=weight_col, verbose=verbose) if normalize not in ("target", "pool", "none"): raise ValueError( f"normalize must be one of 'target', 'pool', 'none'; got {normalize!r}." ) self.match_variance = match_variance self.normalize = normalize self.max_iter = max_iter self.tol = tol self.ridge = ridge self.limit_penalty = limit_penalty md = self.matching_data if md.has_aggregate_target and match_variance: no_mean = [ f for f, stats in md.aggregate_target.numeric.items() if "std" in stats and "mean" not in stats ] if no_mean: raise ValueError( f"Cannot balance the variance of {no_mean}: the target discloses " "a std but no mean, so the variance is taken around the weighted " "mean, which is not a moment constraint reweighting can solve. " "Disclose the mean too, or use AggregateConstraintSatisfactionMatcher." ) # Fitted preprocessor plus pool features and target mean/std in a single # consistent output feature space. An AggregateTarget needs the # calculator that dichotomizes quantile-disclosed features and says # which output columns the target actually discloses. objective = "aggregate_beta" if md.has_aggregate_target else "beta" self.balance_calculator = BalanceCalculator(md, objective) self.preprocessor = self.balance_calculator.preprocessor
[docs] def get_params(self) -> Dict: """Return the weighter's configuration parameters as a dict.""" return { "match_variance": self.match_variance, "normalize": self.normalize, "max_iter": self.max_iter, "tol": self.tol, "ridge": self.ridge, "limit_penalty": self.limit_penalty, }
def _limit_columns(self) -> set: """Output columns holding the indicator of a disclosed min / max.""" md = self.matching_data if not md.has_aggregate_target: return set() out_features = self.preprocessor.output_headers["all"] columns = set() for feature, stats in md.aggregate_target.numeric.items(): for q, _ in stats.get("quantile", []): if q in (0.0, 1.0): prefix = f"{feature}_q{q}" columns |= {c for c in out_features if c.startswith(prefix + "_")} return columns def _numeric_features_with_disclosed_std(self) -> List[str]: md = self.matching_data if md.has_aggregate_target: return [ f for f in md.aggregate_target.headers.numeric if "std" in md.aggregate_target.numeric[f] and "mean" in md.aggregate_target.numeric[f] ] return list(md.headers.numeric) def _build_constraints(self) -> Tuple[np.ndarray, List[str], np.ndarray]: """ Build the (n_pool, n_constraints) centered-and-scaled constraint matrix used to solve for weights, a human-readable label per constraint column (for diagnostics/reporting) and the per-constraint soft-constraint penalty (0 for an exact constraint). """ pool = self.balance_calculator.pool.cpu().numpy() target_mean = self.balance_calculator.target_mean.cpu().numpy().reshape(-1) target_std = self.balance_calculator.target_std.cpu().numpy().reshape(-1) pool_std = pool.std(axis=0) out_features = self.preprocessor.output_headers["all"] n_pool, n_features = pool.shape constrained = self.balance_calculator.constrained limit_columns = self._limit_columns() columns = [] labels = [] penalty = [] for j in range(n_features): if not constrained[j]: continue penalty.append( self.limit_penalty if out_features[j] in limit_columns else 0.0 ) scale = ( target_std[j] if target_std[j] > 0 else (pool_std[j] if pool_std[j] > 0 else 1.0) ) columns.append((pool[:, j] - target_mean[j]) / scale) labels.append(f"{out_features[j]} (mean)") if self.match_variance: variance_features = set(self._numeric_features_with_disclosed_std()) for j, feature in enumerate(out_features): if feature not in variance_features: continue dev = (pool[:, j] - target_mean[j]) ** 2 target_var = target_std[j] ** 2 # Give the variance term its own scale so it neither dominates # nor is swamped by the mean terms purely due to units (same # concern handled in # AggregateConstraintSatisfactionMatcher._get_target_variance_targets). scale2 = target_var if target_var > 0 else max(dev.std(), 1.0) columns.append((dev - target_var) / scale2) labels.append(f"{feature} (variance)") penalty.append(0.0) Z = np.column_stack(columns) if columns else np.zeros((n_pool, 0)) return Z, labels, np.array(penalty) def _fit(self) -> "EntropyBalanceWeighter": md = self.matching_data n_pool = len(md.get_population(md.pool_name)) n_target = ( md.aggregate_target.n if md.has_aggregate_target else len(md.get_population(md.target_name)) ) Z, labels, penalty = self._build_constraints() self.constraint_labels = labels weights, diagnostics = _solve_entropy_weights( Z, max_iter=self.max_iter, tol=self.tol, ridge=self.ridge, penalty=penalty ) if self.normalize == "target": weights = weights * (n_target / weights.sum()) elif self.normalize == "pool": weights = weights * (n_pool / weights.sum()) self.weights = weights self.diagnostics = diagnostics self.diagnostics["effective_sample_size"] = effective_sample_size(weights) if not diagnostics["converged"]: logger.warning( f"{self.__class__.__name__} did not converge within " f"{self.max_iter} iterations (max constraint violation = " f"{diagnostics['max_constraint_violation']:.4g}). Weights may " "not exactly balance the requested moments; consider " "increasing max_iter/ridge, or check for unmatchable (e.g. " "near-extreme or non-overlapping) covariates." ) elif self.verbose: logger.info( f"{self.__class__.__name__} converged in {diagnostics['n_iter']} " "iterations. Effective sample size: " f"{self.diagnostics['effective_sample_size']:.1f} / {n_pool} pool patients." ) return self
[docs]class MAICWeighter(EntropyBalanceWeighter): """ Matching-Adjusted Indirect Comparison (MAIC; Signorovitch et al., 2010, "Comparative effectiveness without head-to-head trials: a method for matching-adjusted indirect comparisons applied to psoriasis clinical trials"). Reweights patient-level ("IPD") pool data so that its weighted means match a target's aggregate statistics -- typically a comparator trial's published Table 1 -- using the method-of-moments / maximum-entropy weighting scheme of ``EntropyBalanceWeighter``, restricted to first moments only. This is the standard MAIC formulation and is appropriate whenever the comparator discloses only means (and category rates), not variances. Use ``EntropyBalanceWeighter(matching_data, match_variance=True)`` directly if the comparator additionally discloses standard deviations you also want to match. :param matching_data: MatchingData whose pool (IPD) is to be weighted to match the target's (typically aggregate, e.g. ``AggregateTarget``) moments. :param normalize: See ``EntropyBalanceWeighter``. Defaults to "target", i.e. weights sum to the target's sample size, matching common MAIC reporting conventions. :param max_iter: Maximum number of Newton iterations. :param tol: Convergence tolerance on the (max-norm) constraint violation. :param ridge: Ridge regularization added to the Newton step for numerical stability. :param weight_col: Name of the column used to store weights on the MatchingData returned by match(). :param verbose: Whether to log fitting diagnostics. :param limit_penalty: See ``EntropyBalanceWeighter``. """ def __init__( self, matching_data: MatchingData, normalize: str = "target", max_iter: int = 200, tol: float = 1e-8, ridge: float = 1e-8, weight_col: str = "sample_weight", verbose: bool = True, limit_penalty: float = 1e-2, ): super().__init__( matching_data, match_variance=False, normalize=normalize, max_iter=max_iter, tol=tol, ridge=ridge, weight_col=weight_col, verbose=verbose, limit_penalty=limit_penalty, )
[docs] def get_params(self) -> Dict: """Return the weighter's configuration parameters as a dict.""" params = super().get_params() del params["match_variance"] return params
[docs]class IPTWWeighter(BaseWeighter): """ Inverse Probability of Treatment Weighting (IPTW; Rosenbaum & Rubin, 1983; see Austin, 2011, "An Introduction to Propensity Score Methods for Reducing the Effects of Confounding in Observational Studies", for the ATT construction used here). Fits a propensity model ``p(X) = P(target | X)`` that classifies pool vs. target patients on their covariates, then reweights each pool patient by the odds ``p / (1 - p)``. The target population keeps weight 1, so the weighted pool is reweighted onto the target's covariate distribution -- i.e. this is the ATT estimand with the target playing the role of the fixed/reference ("treated") group, which matches the convention used by ``MAICWeighter``/``EntropyBalanceWeighter`` and the package's typical use case of building an external comparator arm that represents a trial population. Unlike ``EntropyBalanceWeighter``, IPTW does not guarantee exact balance on any particular moment: it only guarantees balance asymptotically, and only if the propensity model is correctly specified. It also requires a patient-level target (there is no pool-vs-target classification problem to fit against a published aggregate Table 1). Its main practical advantages over entropy balancing are that it scales to many covariates without requiring the target's moments to be inside the pool's convex hull, and that it is the most widely recognized/reported method in the observational literature. :param matching_data: MatchingData whose pool is to be weighted. The target must be patient-level (not an ``AggregateTarget``). :param classifier: A fitted-or-unfitted sklearn-compatible classifier exposing ``predict_proba``, used to estimate ``P(target | X)``. It is cloned before fitting, so passing a pre-fitted instance does not reuse its fit. Defaults to ``LogisticRegression(max_iter=1000)``, the standard choice for propensity score estimation. :param trim_quantiles: Optional ``(low, high)`` quantiles (e.g. ``(0.01, 0.99)``) at which to clip the fitted weights. Extreme weights (driven by pool patients whose covariates make them look almost certainly pool or almost certainly target) are the main practical failure mode of IPTW; trimming trades a little bias for a large reduction in variance. Left unset (``None``) by default so the raw weights are returned untouched. :param weight_col: Name of the column used to store weights on the MatchingData returned by match(). :param verbose: Whether to log fitting diagnostics. """ def __init__( self, matching_data: MatchingData, classifier: Optional[BaseEstimator] = None, trim_quantiles: Optional[Tuple[float, float]] = None, weight_col: str = "sample_weight", verbose: bool = True, ): super().__init__(matching_data, weight_col=weight_col, verbose=verbose) if self.matching_data.has_aggregate_target: raise ValueError( "IPTWWeighter requires a patient-level target (it fits a " "pool-vs-target propensity model), so it cannot be used with " "an AggregateTarget. Use EntropyBalanceWeighter/MAICWeighter " "for aggregate (e.g. published Table 1) targets." ) if trim_quantiles is not None: lo, hi = trim_quantiles if not (0 <= lo < hi <= 1): raise ValueError( f"trim_quantiles must satisfy 0 <= low < high <= 1; got {trim_quantiles!r}." ) self.classifier = classifier self.trim_quantiles = trim_quantiles # Reuse BetaBalance purely to get a fitted preprocessor plus pool and # target feature tensors in a single consistent (one-hot categoric + # numeric passthrough) output space, as EntropyBalanceWeighter does. self.balance_calculator = BalanceCalculator(self.matching_data, "beta") self.preprocessor = self.balance_calculator.preprocessor
[docs] def get_params(self) -> Dict: """Return the weighter's configuration parameters as a dict.""" return { "classifier": self.classifier, "trim_quantiles": self.trim_quantiles, }
def _fit(self) -> "IPTWWeighter": pool = self.balance_calculator.pool.cpu().numpy() target = self.balance_calculator.target.cpu().numpy() n_pool = len(pool) X = np.vstack([pool, target]) y = np.concatenate([np.zeros(n_pool), np.ones(len(target))]) scaler = StandardScaler() X = scaler.fit_transform(X) clf = ( clone(self.classifier) if self.classifier is not None else LogisticRegression(max_iter=1000) ) clf.fit(X, y) propensity_score = clf.predict_proba(X[:n_pool])[:, 1] # Clip away from 0/1 to avoid infinite weights from a # (near-)perfectly separable model. propensity_score = np.clip(propensity_score, 1e-4, 1 - 1e-4) weights = propensity_score / (1 - propensity_score) # Also kept (unclipped) so plot_iptw_propensity_distributions() can show # where the target itself sits on the fitted model. target_propensity_score = clf.predict_proba(X[n_pool:])[:, 1] n_trimmed = 0 if self.trim_quantiles is not None: lo, hi = np.quantile(weights, self.trim_quantiles) n_trimmed = int(np.sum((weights < lo) | (weights > hi))) weights = np.clip(weights, lo, hi) self.propensity_model = clf self.propensity_score = propensity_score self.target_propensity_score = target_propensity_score self.weights = weights self.diagnostics = { "effective_sample_size": effective_sample_size(weights), "n_trimmed": n_trimmed, } if self.verbose: logger.info( f"{self.__class__.__name__} fit {str(clf).split('(')[0]}. " f"Effective sample size: {self.diagnostics['effective_sample_size']:.1f} " f"/ {n_pool} pool patients." + (f" Trimmed {n_trimmed} extreme weights." if n_trimmed else "") ) return self
[docs]def weighted_balance_table(weighter: BaseWeighter) -> pd.DataFrame: """ Return a table comparing each balanced feature's weighted (and, for reference, unweighted) pool moment to the target moment -- a quick post-hoc check of how well match() balanced the pool. For EntropyBalanceWeighter/MAICWeighter, residuals should be ~0 for every constrained row; for IPTWWeighter, which does not solve for exact balance, this is instead a diagnostic of how much balance improved relative to the unweighted pool. :param weighter: A fitted EntropyBalanceWeighter/MAICWeighter/IPTWWeighter, i.e. one on which match() has already been called. """ _check_fitted(weighter) pool = weighter.balance_calculator.pool.cpu().numpy() target_mean = weighter.balance_calculator.target_mean.cpu().numpy().reshape(-1) target_std = weighter.balance_calculator.target_std.cpu().numpy().reshape(-1) out_features = weighter.preprocessor.output_headers["all"] w = weighter.weights variance_features = ( set(weighter._numeric_features_with_disclosed_std()) if getattr(weighter, "match_variance", False) else set() ) constrained = weighter.balance_calculator.constrained disclosed = weighter.balance_calculator.disclosed rows = [] for j, feature in enumerate(out_features): if disclosed[j] and not constrained[j]: continue # a level implied by the others rows.append( { "feature": feature, "moment": "mean", # undisclosed: nothing to compare to, but still worth seeing move "target": target_mean[j] if disclosed[j] else np.nan, "unweighted_pool": pool[:, j].mean(), "weighted_pool": np.average(pool[:, j], weights=w), } ) if feature in variance_features: rows.append( { "feature": feature, "moment": "variance", "target": target_std[j] ** 2, "unweighted_pool": pool[:, j].var(), "weighted_pool": np.average( (pool[:, j] - np.average(pool[:, j], weights=w)) ** 2, weights=w, ), } ) return pd.DataFrame(rows)
def _check_is_iptw(weighter: BaseWeighter) -> None: if not hasattr(weighter, "propensity_score"): raise TypeError( "plot_iptw_propensity_distributions() requires an IPTWWeighter " f"(got {type(weighter).__name__}), since only IPTWWeighter fits a " "propensity model -- MAICWeighter/EntropyBalanceWeighter solve " "directly for balancing weights without one." )
[docs]def plot_iptw_propensity_distributions(weighter: "IPTWWeighter"): """ Plot histograms of the estimated propensity score for the pool and target populations, before vs. after IPTW weighting -- the weighting analogue of ``pybalance.propensity.plot_propensity_score_match_distributions``. Unlike a Matcher, a Weighter never drops patients, so there is no matched subset to compare against; instead, the "before" panel shows every pool patient counted equally and the "after" panel shows the same propensity scores counted by their fitted IPTW weight (the target always keeps weight 1). A successful fit should show the "after" pool histogram move towards the target's. :param weighter: A fitted ``IPTWWeighter``, i.e. one on which match() has already been called. """ _check_fitted(weighter) _check_is_iptw(weighter) md = weighter.matching_data pool_name, target_name = md.pool_name, md.target_name data = pd.concat( [ pd.DataFrame.from_dict( { "propensity": weighter.propensity_score, "weight": np.ones_like(weighter.propensity_score), "weighted": False, "population": pool_name, } ), pd.DataFrame.from_dict( { "propensity": weighter.target_propensity_score, "weight": np.ones_like(weighter.target_propensity_score), "weighted": False, "population": target_name, } ), pd.DataFrame.from_dict( { "propensity": weighter.propensity_score, "weight": weighter.weights, "weighted": True, "population": pool_name, } ), pd.DataFrame.from_dict( { "propensity": weighter.target_propensity_score, "weight": np.ones_like(weighter.target_propensity_score), "weighted": True, "population": target_name, } ), ] ) g = sns.FacetGrid( data=data, col="weighted", col_order=[False, True], height=4, xlim=[0, 1] ) g.map_dataframe( sns.histplot, bins=24, binrange=(0, 1), x="propensity", weights="weight", hue="population", hue_order=[pool_name, target_name], alpha=0.5, common_norm=False, stat="probability", ) [ax.grid(True) for axes in g.axes for ax in axes] legend_patches = [ matplotlib.patches.Patch(color=sns.color_palette()[0], label=pool_name), matplotlib.patches.Patch(color=sns.color_palette()[1], label=target_name), ] plt.legend(handles=legend_patches) return g