pymte 1.0.0__py3-none-any.whl

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.
pymte/__init__.py ADDED
@@ -0,0 +1,74 @@
1
+ """Instrumental variables: extrapolation by marginal treatment effects.
2
+
3
+ This package is a Python port of the R package ``ivmte`` by Joshua Shea and
4
+ Alexander Torgovitsky. It implements the moment-based marginal treatment
5
+ effect framework of Mogstad, Santos and Torgovitsky (2018) for point and
6
+ partial identification of treatment parameters. The modules and functions
7
+ follow the layout and the names of the R package.
8
+ """
9
+
10
+ from pymte.audit import audit, rhalton
11
+ from pymte.datasets import load_ae, load_sim_data
12
+ from pymte.design import design
13
+ from pymte.ivlike import iv_estimate
14
+ from pymte.lp import (
15
+ bound,
16
+ criterion_min,
17
+ lp_setup,
18
+ lp_setup_bound,
19
+ lp_setup_criterion,
20
+ lp_setup_criterion_boot,
21
+ qp_setup,
22
+ qp_setup_bound,
23
+ qp_setup_criterion,
24
+ )
25
+ from pymte.mst import (
26
+ IVMTEResult,
27
+ bound_ci,
28
+ bound_pvalue,
29
+ gen_s_set,
30
+ gen_target,
31
+ gmm_estimate,
32
+ ivmte,
33
+ ivmte_estimate,
34
+ )
35
+ from pymte.mtr import MTRSpec, gen_gamma, polyparse
36
+ from pymte.plots import plot_mte, plot_mtr, plot_weights
37
+ from pymte.propensity import Propensity, propensity
38
+ from pymte.splines import USpline
39
+
40
+ __all__ = [
41
+ "IVMTEResult",
42
+ "MTRSpec",
43
+ "Propensity",
44
+ "USpline",
45
+ "audit",
46
+ "bound",
47
+ "bound_ci",
48
+ "bound_pvalue",
49
+ "criterion_min",
50
+ "design",
51
+ "gen_gamma",
52
+ "gen_s_set",
53
+ "gen_target",
54
+ "gmm_estimate",
55
+ "iv_estimate",
56
+ "ivmte",
57
+ "ivmte_estimate",
58
+ "load_ae",
59
+ "load_sim_data",
60
+ "lp_setup",
61
+ "lp_setup_bound",
62
+ "lp_setup_criterion",
63
+ "lp_setup_criterion_boot",
64
+ "plot_mte",
65
+ "plot_mtr",
66
+ "plot_weights",
67
+ "polyparse",
68
+ "propensity",
69
+ "qp_setup",
70
+ "qp_setup_bound",
71
+ "qp_setup_criterion",
72
+ "rhalton",
73
+ ]
74
+ __version__ = "1.0.0"
pymte/audit.py ADDED
@@ -0,0 +1,547 @@
1
+ """The audit procedure for partial identification.
2
+
3
+ Shape restrictions are first imposed on a small initial grid. After solving
4
+ for the bounds, the restrictions are checked on the finer audit grid; grid
5
+ points where either bounding solution violates a restriction are added to
6
+ the constraint set and the problem is solved again, until no violations
7
+ remain or ``audit_max`` rounds have been performed. The grids, the rule for
8
+ choosing which violations to add and the termination rules follow the R
9
+ package; the few places where the two differ (the end points of a custom
10
+ ``audit_u``, the full-grid test, and the ranking of violations when
11
+ ``audit_add`` binds) are listed in the migration notes.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import warnings
17
+ from collections.abc import Sequence
18
+ from dataclasses import dataclass, field
19
+ from typing import Any
20
+
21
+ import numpy as np
22
+ import pandas as pd
23
+ from numpy.typing import ArrayLike, NDArray
24
+
25
+ from pymte.lp import (
26
+ STATUS_STRINGS,
27
+ Criterion,
28
+ SolveResult,
29
+ bound,
30
+ criterion_min,
31
+ lp_setup,
32
+ magnitude,
33
+ )
34
+ from pymte.monobound import KINDS, Grids, ShapeConstraints, combinemonobound, genmonobound_a
35
+ from pymte.mtr import MTRSpec
36
+
37
+
38
+ class AuditError(RuntimeError):
39
+ """The criterion or a bound problem could not be solved.
40
+
41
+ Attributes
42
+ ----------
43
+ status : int
44
+ Canonical solver status code of the failed problem.
45
+ stage : {"criterion", "bound"}
46
+ Which problem failed.
47
+ grids : Grids or None
48
+ The grids in use, so that a retry can keep the audit grid.
49
+ """
50
+
51
+ def __init__(
52
+ self, message: str, status: int, stage: str = "criterion", grids: Grids | None = None
53
+ ) -> None:
54
+ super().__init__(message)
55
+ self.status = status
56
+ self.stage = stage
57
+ self.grids = grids
58
+
59
+
60
+ @dataclass
61
+ class AuditResult:
62
+ """Outcome of the audit procedure.
63
+
64
+ Attributes
65
+ ----------
66
+ lower, upper : float
67
+ Bounds on the target parameter.
68
+ theta_min, theta_max, theta_crit : numpy.ndarray
69
+ Stacked MTR coefficients at the lower bound, the upper bound and the
70
+ criterion minimum of the final round.
71
+ criterion : float
72
+ Minimum criterion in the final round.
73
+ audit_count : int
74
+ Number of rounds performed.
75
+ violations : pandas.DataFrame
76
+ Violations found on the audit grid in the final round (empty when
77
+ the audit finished cleanly).
78
+ constraints : ShapeConstraints
79
+ Shape restrictions imposed in the final round.
80
+ grids : Grids
81
+ The audit grid and the initial constraint grid that were used.
82
+ status : dict
83
+ Solver status codes for ``"criterion"``, ``"min"`` and ``"max"``.
84
+ runtime : dict
85
+ Solver time in seconds for the same three problems (final round).
86
+ messages : list of str
87
+ Progress log, in the wording of the R package.
88
+ """
89
+
90
+ lower: float
91
+ upper: float
92
+ theta_min: NDArray[np.float64]
93
+ theta_max: NDArray[np.float64]
94
+ theta_crit: NDArray[np.float64]
95
+ criterion: float
96
+ audit_count: int
97
+ violations: pd.DataFrame
98
+ constraints: ShapeConstraints
99
+ grids: Grids
100
+ status: dict[str, int]
101
+ runtime: dict[str, float]
102
+ messages: list[str] = field(default_factory=list)
103
+
104
+
105
+ def status_string(status: int) -> str:
106
+ """Describe a canonical solver status code."""
107
+ return STATUS_STRINGS.get(status, "unknown")
108
+
109
+
110
+ def fmt_result(x: float) -> str:
111
+ """Format numbers the way ``print.ivmte`` does."""
112
+ if x == 0:
113
+ return "0"
114
+ if abs(x) < 1:
115
+ return f"{x:.7g}"
116
+ if abs(x) < 1e7:
117
+ return f"{round(x, 4):.7g}"
118
+ return f"{x:.7e}"
119
+
120
+
121
+ # -- grids --------------------------------------------------------------------
122
+
123
+
124
+ def rhalton(n: int, base: int = 2) -> NDArray[np.float64]:
125
+ """First ``n`` points of the Halton (van der Corput) sequence, as in R ``rhalton``."""
126
+ out = np.empty(n)
127
+ for j in range(1, n + 1):
128
+ f, r, i = 1.0, 0.0, j
129
+ while i > 0:
130
+ f /= base
131
+ r += f * (i % base)
132
+ i //= base
133
+ out[j - 1] = r
134
+ return out
135
+
136
+
137
+ def _u_grid(n: int) -> NDArray[np.float64]:
138
+ """Sorted grid ``{0, 1} + rhalton(n)`` rounded to 8 decimals, as in the R package."""
139
+ if n <= 0:
140
+ return np.array([0.0, 1.0])
141
+ return np.sort(np.concatenate([[0.0, 1.0], np.round(rhalton(n), 8)]))
142
+
143
+
144
+ def _gen_grids(
145
+ data: pd.DataFrame,
146
+ xvars: Sequence[str],
147
+ *,
148
+ initgrid_nx: int,
149
+ initgrid_nu: int,
150
+ audit_nx: int,
151
+ audit_nu: int,
152
+ initgrid_x: pd.DataFrame | None,
153
+ initgrid_u: ArrayLike | None,
154
+ audit_x: pd.DataFrame | None,
155
+ audit_u: ArrayLike | None,
156
+ rng: np.random.Generator,
157
+ ) -> Grids:
158
+ """Sample the audit grid and the initial constraint grid as the R package does.
159
+
160
+ The audit grid in ``x`` is a uniform sample (capped by the support size)
161
+ of the distinct covariate rows, the initial grid a subsample of it; the
162
+ ``u`` grids are Halton points plus the end points. Explicit grids
163
+ override the sampled ones.
164
+ """
165
+ xvars = list(xvars)
166
+ if not xvars:
167
+ support = pd.DataFrame(index=pd.RangeIndex(1))
168
+ init_index = np.array([0])
169
+ else:
170
+ if audit_x is None:
171
+ full = data[xvars].drop_duplicates().reset_index(drop=True)
172
+ take = min(audit_nx, len(full))
173
+ support = full.iloc[np.sort(rng.choice(len(full), take, replace=False))]
174
+ support = support.reset_index(drop=True)
175
+ else:
176
+ support = audit_x[xvars].reset_index(drop=True)
177
+ if initgrid_x is None:
178
+ take = min(initgrid_nx, len(support))
179
+ init_index = np.sort(rng.choice(len(support), take, replace=False))
180
+ else:
181
+ init = initgrid_x[xvars].reset_index(drop=True)
182
+ combined = pd.concat([init, support], ignore_index=True)
183
+ is_init = np.arange(len(combined)) < len(init)
184
+ keep = ~(combined.duplicated().to_numpy() & ~is_init)
185
+ support = combined.loc[keep].reset_index(drop=True)
186
+ init_index = np.flatnonzero(is_init[keep])
187
+ if audit_u is None:
188
+ a_u = _u_grid(audit_nu)
189
+ else:
190
+ a_u = np.asarray(audit_u, dtype=float)
191
+ if initgrid_u is not None:
192
+ a_u = np.union1d(a_u, np.asarray(initgrid_u, dtype=float))
193
+ a_u = np.union1d(a_u, [0.0, 1.0])
194
+ if initgrid_u is not None:
195
+ i_u = np.union1d(np.asarray(initgrid_u, dtype=float), [0.0, 1.0])
196
+ elif initgrid_nu <= 0:
197
+ i_u = np.array([0.0, 1.0])
198
+ elif audit_u is None:
199
+ i_u = _u_grid(initgrid_nu)
200
+ else:
201
+ take = min(len(a_u), initgrid_nu)
202
+ i_u = np.sort(rng.choice(a_u, take, replace=False))
203
+ return Grids(support, a_u, np.asarray(init_index, dtype=np.intp), i_u)
204
+
205
+
206
+ # -- violations ---------------------------------------------------------------
207
+
208
+
209
+ def _violations(
210
+ cons: ShapeConstraints, thetas: Sequence[NDArray[np.float64]], tol: float
211
+ ) -> pd.DataFrame:
212
+ """Constraint rows violated by any of the candidate solutions.
213
+
214
+ Returns a frame with columns ``row`` (position in ``cons``), ``kind``,
215
+ ``x_index``, ``u`` and ``diff`` (largest residual across the solutions).
216
+ """
217
+ diff = np.max(np.column_stack([cons.residuals(t) for t in thetas]), axis=1)
218
+ pos = np.flatnonzero(diff > tol)
219
+ return pd.DataFrame(
220
+ {
221
+ "row": pos,
222
+ "kind": cons.kind[pos],
223
+ "x_index": cons.x_index[pos],
224
+ "u": cons.u[pos],
225
+ "diff": diff[pos],
226
+ }
227
+ )
228
+
229
+
230
+ def select_violations(viol: pd.DataFrame, audit_add: int) -> pd.DataFrame:
231
+ """Choose which violated points to add, following the R package.
232
+
233
+ All violations are added when there are at most ``audit_add``. Otherwise
234
+ the worst violation of every (restriction, covariate cell) group is
235
+ taken first, then the second worst of every group, and so on until at
236
+ least ``audit_add`` points are selected.
237
+ """
238
+ if len(viol) <= audit_add:
239
+ return viol
240
+ kind_order = {k: i for i, k in enumerate(KINDS)}
241
+ v = viol.assign(_k=viol["kind"].map(kind_order))
242
+ v = v.sort_values(["_k", "x_index", "diff"], ascending=[True, True, False])
243
+ v["rank"] = v.groupby(["_k", "x_index"]).cumcount() + 1
244
+ counts = v["rank"].value_counts().sort_index().cumsum()
245
+ if counts.iloc[0] >= audit_add:
246
+ chosen = v[v["rank"] == 1]
247
+ else:
248
+ k = int(counts.index[int(np.searchsorted(counts.to_numpy(), audit_add))])
249
+ full = v[v["rank"] <= k - 1]
250
+ extra = v[v["rank"] == k].sort_values("diff", ascending=False)
251
+ chosen = pd.concat([full, extra.head(audit_add - len(full))])
252
+ chosen = chosen.sort_values(["_k", "x_index", "diff"], ascending=[True, True, False])
253
+ return chosen.drop(columns=["_k", "rank"])
254
+
255
+
256
+ # -- the audit loop -----------------------------------------------------------
257
+
258
+
259
+ def _relaxed_tol(tol: float) -> float:
260
+ # R: (tol / 10^magnitude) * 10^(magnitude / 2), e.g. 1e-6 -> 1e-3.
261
+ mag = magnitude(tol)
262
+ if mag is None:
263
+ return tol
264
+ return float((tol / 10**mag) * 10 ** (mag / 2))
265
+
266
+
267
+ _DEFAULT_NOTE = {
268
+ "m0.lb": "min. observed outcome by default",
269
+ "m1.lb": "min. observed outcome by default",
270
+ "m0.ub": "max. observed outcome by default",
271
+ "m1.ub": "max. observed outcome by default",
272
+ }
273
+
274
+
275
+ def _describe_restriction(kind: str, restrictions: dict[str, Any], defaults: Sequence[str]) -> str:
276
+ key = kind.replace(".", "_")
277
+ value = restrictions[key]
278
+ if isinstance(value, bool):
279
+ return f"{key} = {value}"
280
+ text = f"{key} = {round(float(value), 6)}"
281
+ if key in defaults:
282
+ text += f" ({_DEFAULT_NOTE[kind]})"
283
+ return text
284
+
285
+
286
+ def _infeasibility_message(
287
+ status: int,
288
+ crit: Criterion,
289
+ equal: NDArray[np.float64] | None,
290
+ cons: ShapeConstraints,
291
+ restrictions: dict[str, Any],
292
+ defaults: Sequence[str],
293
+ audit_tol: float,
294
+ solver: str | None,
295
+ options: dict[str, Any] | None,
296
+ ) -> str:
297
+ """Diagnose an infeasible criterion problem as the R package does.
298
+
299
+ The criterion is minimised again without the shape restrictions and the
300
+ restrictions that solution violates are named, since incoherent shape
301
+ restrictions are the likely cause of an empty parameter space.
302
+ """
303
+ proved = "infeasible" if status == 2 else "infeasible or unbounded"
304
+ message = f"No solution since the solver proved the model was {proved}."
305
+ _, theta, _ = criterion_min(crit, lp_setup(crit, None, equal), solver, options)
306
+ if theta is not None:
307
+ violated = cons.residuals(theta) > audit_tol
308
+ kinds = [k for k in KINDS if k in set(cons.kind[violated])]
309
+ if kinds:
310
+ named = ", ".join(_describe_restriction(k, restrictions, defaults) for k in kinds)
311
+ message += (
312
+ " The model should only be infeasible if the implied parameter space is "
313
+ "empty. The likely cause of an empty parameter space is incoherent shape "
314
+ f"restrictions. For example, {named} are all set simultaneously. Try "
315
+ "changing the shape constraints on the MTR functions."
316
+ )
317
+ if status == 3:
318
+ message += (
319
+ " The model may be unbounded if the initial grid is too small. Try increasing "
320
+ "the parameters 'initgrid_nx' and 'initgrid_nu'."
321
+ )
322
+ return message
323
+
324
+
325
+ def _check_criterion(res: SolveResult, alt: str) -> None:
326
+ if res.status == 4:
327
+ raise AuditError(
328
+ "No solution to minimizing the criterion since the model is unbounded. " + alt,
329
+ res.status,
330
+ )
331
+ if res.status == 5:
332
+ raise AuditError(
333
+ "No solution to minimizing the criterion due to numerical issues. " + alt, res.status
334
+ )
335
+ if res.x is None:
336
+ raise AuditError(
337
+ "No solution provided by the solver when minimizing the criterion. " + alt,
338
+ res.status,
339
+ )
340
+
341
+
342
+ def audit(
343
+ spec0: MTRSpec,
344
+ spec1: MTRSpec,
345
+ crit: Criterion,
346
+ gstar: NDArray[np.float64],
347
+ data: pd.DataFrame,
348
+ restrictions: dict[str, Any],
349
+ *,
350
+ defaults: Sequence[str] = (),
351
+ equal: NDArray[np.float64] | None = None,
352
+ criterion_tol: float = 1e-4,
353
+ audit_tol: float = 1e-6,
354
+ audit_add: int = 100,
355
+ audit_max: int = 25,
356
+ initgrid_nx: int = 20,
357
+ initgrid_nu: int = 20,
358
+ audit_nx: int = 2500,
359
+ audit_nu: int = 25,
360
+ initgrid_x: pd.DataFrame | None = None,
361
+ initgrid_u: ArrayLike | None = None,
362
+ audit_x: pd.DataFrame | None = None,
363
+ audit_u: ArrayLike | None = None,
364
+ audit_grid: Grids | None = None,
365
+ rng: np.random.Generator | None = None,
366
+ solver: str | None = None,
367
+ solver_options_criterion: dict[str, Any] | None = None,
368
+ solver_options_bounds: dict[str, Any] | None = None,
369
+ log: list[str] | None = None,
370
+ ) -> AuditResult:
371
+ """Compute bounds on the target parameter with the audit procedure.
372
+
373
+ Parameters
374
+ ----------
375
+ spec0, spec1 : MTRSpec
376
+ MTR specifications.
377
+ crit : L1Criterion or LSCriterion
378
+ Criterion measuring the fit to the data.
379
+ gstar : numpy.ndarray
380
+ Stacked target coefficients ``(gstar0, gstar1)``.
381
+ data : pandas.DataFrame
382
+ Estimation sample, from which the covariate grids are drawn.
383
+ restrictions : dict
384
+ Shape restrictions, see :func:`pymte.monobound.genmonobound_a`.
385
+ defaults : sequence of str, optional
386
+ Keys of ``restrictions`` that were not set by the user but taken
387
+ from the range of the outcome; named as such in error messages.
388
+ equal : numpy.ndarray, optional
389
+ Equality rows on the coefficients.
390
+ criterion_tol : float
391
+ Relative tolerance on the criterion in the bound problems.
392
+ audit_tol : float
393
+ Violations below this size are ignored.
394
+ audit_add : int
395
+ Maximum number of violated points added per round (see
396
+ :func:`select_violations`).
397
+ audit_max : int
398
+ Maximum number of rounds.
399
+ initgrid_nx, initgrid_nu, audit_nx, audit_nu : int
400
+ Sizes of the initial constraint grid and the audit grid.
401
+ initgrid_x, initgrid_u, audit_x, audit_u : optional
402
+ Explicit grids overriding the sampled ones. The end points 0 and 1
403
+ are always added to the ``u`` grids (R adds them to a custom
404
+ ``audit_u`` only when ``initgrid_u`` is given as well).
405
+ audit_grid : Grids, optional
406
+ Reuse these grids instead of drawing new ones (bootstrap replicates).
407
+ rng : numpy.random.Generator, optional
408
+ Source of randomness for sampling covariate rows.
409
+ solver : str, optional
410
+ Solver name, see :mod:`pymte.lp`.
411
+ solver_options_criterion, solver_options_bounds : dict, optional
412
+ Options for the criterion and the bound problems.
413
+ log : list of str, optional
414
+ Progress messages are appended here.
415
+
416
+ Returns
417
+ -------
418
+ AuditResult
419
+
420
+ Raises
421
+ ------
422
+ AuditError
423
+ When the criterion or a bound problem has no usable solution. An
424
+ infeasible criterion problem is diagnosed by solving it again
425
+ without the shape restrictions and naming the restrictions that
426
+ solution violates.
427
+ """
428
+ log = log if log is not None else []
429
+ alt = "Try relaxing 'criterion_tol' or the shape restrictions."
430
+ if audit_grid is None:
431
+ xvars = sorted(set(spec0.covariates) | set(spec1.covariates))
432
+ audit_grid = _gen_grids(
433
+ data, xvars, initgrid_nx=initgrid_nx, initgrid_nu=initgrid_nu, audit_nx=audit_nx,
434
+ audit_nu=audit_nu, initgrid_x=initgrid_x, initgrid_u=initgrid_u, audit_x=audit_x,
435
+ audit_u=audit_u, rng=rng or np.random.default_rng(),
436
+ ) # fmt: skip
437
+ grids = audit_grid
438
+ support = grids.support
439
+ all_x = np.arange(len(support))
440
+ current = genmonobound_a(spec0, spec1, support, grids.init_index, grids.init_u, restrictions)
441
+ audit_cons = genmonobound_a(spec0, spec1, support, all_x, grids.audit_u, restrictions)
442
+ full_grid = len(grids.init_index) == len(support) and len(grids.init_u) == len(grids.audit_u)
443
+ log.append(" Generating initial constraint grid...")
444
+
445
+ prev: pd.DataFrame | None = None
446
+ same = 0
447
+ count = 1
448
+ while True:
449
+ log.append(f"\n Audit count: {count}")
450
+ model = lp_setup(crit, current, equal)
451
+ crit_res, theta_crit, crit_min = criterion_min(
452
+ crit, model, solver, solver_options_criterion
453
+ )
454
+ if crit_res.status in (2, 3):
455
+ raise AuditError(
456
+ _infeasibility_message(
457
+ crit_res.status, crit, equal, current, restrictions, defaults, audit_tol,
458
+ solver, solver_options_criterion,
459
+ ),
460
+ crit_res.status,
461
+ "criterion",
462
+ grids,
463
+ ) # fmt: skip
464
+ _check_criterion(crit_res, alt)
465
+ assert theta_crit is not None and crit_min is not None
466
+ log.append(f" Minimum criterion: {fmt_result(crit_min)}")
467
+ log.append(" Obtaining bounds...")
468
+ (min_res, theta_min), (max_res, theta_max) = bound(
469
+ crit, model, gstar, crit_min, criterion_tol, solver, solver_options_bounds
470
+ )
471
+ if theta_min is None or theta_max is None or min_res.obj is None or max_res.obj is None:
472
+ bad = min_res if theta_min is None else max_res
473
+ raise AuditError(
474
+ f"The {'minimization' if bad is min_res else 'maximization'} problem for the "
475
+ f"bounds returned no solution (status: {bad.status_str}). {alt}",
476
+ bad.status,
477
+ "bound",
478
+ grids,
479
+ )
480
+ result = AuditResult(
481
+ lower=float(min_res.obj),
482
+ upper=float(max_res.obj),
483
+ theta_min=theta_min,
484
+ theta_max=theta_max,
485
+ theta_crit=theta_crit,
486
+ criterion=float(crit_min),
487
+ audit_count=count,
488
+ violations=pd.DataFrame(),
489
+ constraints=current,
490
+ grids=grids,
491
+ status={"criterion": crit_res.status, "min": min_res.status, "max": max_res.status},
492
+ runtime={
493
+ "criterion": crit_res.runtime,
494
+ "min": min_res.runtime,
495
+ "max": max_res.runtime,
496
+ },
497
+ messages=log,
498
+ )
499
+ viol = _violations(audit_cons, [theta_min, theta_max], audit_tol)
500
+ if full_grid and len(viol):
501
+ # Nothing can be added; the violations are solver precision.
502
+ new_tol = _relaxed_tol(audit_tol)
503
+ worst = viol["diff"].max()
504
+ viol = viol[viol["diff"] > new_tol]
505
+ warnings.warn(
506
+ f"Violations of the shape constraints were found although the initial grid "
507
+ f"equals the audit grid. The audit tolerance was raised from {audit_tol:g} "
508
+ f"to {new_tol:g}; the largest violation was {worst:.3g}.",
509
+ stacklevel=2,
510
+ )
511
+ result.violations = viol
512
+ break
513
+ if prev is not None and len(viol) and len(prev) == len(viol):
514
+ same = (
515
+ same + 1 if viol.reset_index(drop=True).equals(prev.reset_index(drop=True)) else 0
516
+ )
517
+ prev = viol
518
+ if same >= 2:
519
+ warnings.warn(
520
+ "Audit is unable to resolve violations: the same set of violations have "
521
+ "persisted for three iterations. This can occur if 'audit_tol' differs from "
522
+ "the tolerance of the solver. Audit is terminated.",
523
+ stacklevel=2,
524
+ )
525
+ result.violations = viol
526
+ break
527
+ if len(viol) == 0:
528
+ log.append(" Violations: 0")
529
+ log.append(" Audit finished.\n")
530
+ break
531
+ log.append(f" Violations: {len(viol)}")
532
+ if count == audit_max:
533
+ warnings.warn(
534
+ f"Audit finished: maximum number of audits (audit_max = {audit_max}) reached. "
535
+ "Try increasing audit_max.",
536
+ stacklevel=2,
537
+ )
538
+ result.violations = viol
539
+ break
540
+ chosen = select_violations(viol, audit_add)
541
+ log.append(f" Expanding constraint grid to include {len(chosen)} additional points...")
542
+ current = combinemonobound(current, audit_cons.subset(chosen["row"].to_numpy()))
543
+ count += 1
544
+ log.append(
545
+ f"Bounds on the target parameter: [{fmt_result(result.lower)}, {fmt_result(result.upper)}]"
546
+ )
547
+ return result
pymte/callcheck.py ADDED
@@ -0,0 +1,50 @@
1
+ """Helpers for checking and dissecting the arguments of :func:`pymte.ivmte`."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Sequence
6
+
7
+ import pandas as pd
8
+ from formulaic import Formula
9
+
10
+
11
+ def get_xz(formula: str) -> tuple[str, str, str | None]:
12
+ """Split an IV-like formula ``y ~ x | z`` into its outcome, regressors and instruments.
13
+
14
+ Parameters
15
+ ----------
16
+ formula : str
17
+ Two-sided formula; the part after ``|`` (optional) lists the
18
+ instruments.
19
+
20
+ Returns
21
+ -------
22
+ tuple
23
+ Outcome name, regressor specification and instrument specification
24
+ (``None`` for OLS).
25
+ """
26
+ lhs, sep, rhs = formula.partition("~")
27
+ if not sep or not lhs.strip():
28
+ raise ValueError(f"IV-like formulas need an outcome on the left-hand side: {formula!r}")
29
+ x_part, bar, z_part = rhs.partition("|")
30
+ return lhs.strip(), x_part.strip(), (z_part.strip() if bar else None)
31
+
32
+
33
+ def formula_vars(formula: str) -> set[str]:
34
+ """Variables referenced by a formula (both sides, instruments included)."""
35
+ vars_: set[str] = set()
36
+ for part in formula.replace("|", "+").split("~"):
37
+ part = part.strip()
38
+ if part:
39
+ vars_ |= set(Formula(part).required_variables)
40
+ return vars_
41
+
42
+
43
+ def required_columns(
44
+ data: pd.DataFrame, formulas: Sequence[str], extra: Sequence[str]
45
+ ) -> list[str]:
46
+ """List the data columns referenced by the formulas; unknown names such as ``u`` are ignored."""
47
+ names: set[str] = set(extra)
48
+ for f in formulas:
49
+ names |= formula_vars(f)
50
+ return [c for c in data.columns if c in names]