treecf 0.2.1__tar.gz → 0.2.2__tar.gz

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.
Files changed (62) hide show
  1. {treecf-0.2.1 → treecf-0.2.2}/PKG-INFO +1 -1
  2. {treecf-0.2.1 → treecf-0.2.2}/pyproject.toml +1 -1
  3. {treecf-0.2.1 → treecf-0.2.2}/rust/Cargo.lock +1 -1
  4. {treecf-0.2.1 → treecf-0.2.2}/rust/Cargo.toml +1 -1
  5. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/__init__.py +4 -1
  6. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/api.py +92 -0
  7. treecf-0.2.2/src/treecf/audit.py +525 -0
  8. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/batch.py +64 -15
  9. {treecf-0.2.1 → treecf-0.2.2}/LICENSE +0 -0
  10. {treecf-0.2.1 → treecf-0.2.2}/README.md +0 -0
  11. {treecf-0.2.1 → treecf-0.2.2}/rust/src/cells.rs +0 -0
  12. {treecf-0.2.1 → treecf-0.2.2}/rust/src/constraints.rs +0 -0
  13. {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/domains.rs +0 -0
  14. {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/mod.rs +0 -0
  15. {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/orderpairs.rs +0 -0
  16. {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/propagation.rs +0 -0
  17. {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/search.rs +0 -0
  18. {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/test_support.rs +0 -0
  19. {treecf-0.2.1 → treecf-0.2.2}/rust/src/ga.rs +0 -0
  20. {treecf-0.2.1 → treecf-0.2.2}/rust/src/interrupt.rs +0 -0
  21. {treecf-0.2.1 → treecf-0.2.2}/rust/src/ir.rs +0 -0
  22. {treecf-0.2.1 → treecf-0.2.2}/rust/src/lib.rs +0 -0
  23. {treecf-0.2.1 → treecf-0.2.2}/rust/src/py.rs +0 -0
  24. {treecf-0.2.1 → treecf-0.2.2}/rust/src/regions.rs +0 -0
  25. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/_errors.py +0 -0
  26. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/_json.py +0 -0
  27. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/aim/__init__.py +0 -0
  28. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/aim/cells.py +0 -0
  29. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/__init__.py +0 -0
  30. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_bounds.py +0 -0
  31. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_domains.py +0 -0
  32. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_orderpairs.py +0 -0
  33. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_propagation.py +0 -0
  34. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/exact.py +0 -0
  35. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/exact_rust.py +0 -0
  36. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/genetic.py +0 -0
  37. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/genetic_rust.py +0 -0
  38. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/regions_rust.py +0 -0
  39. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/__init__.py +0 -0
  40. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/compile.py +0 -0
  41. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/flatten.py +0 -0
  42. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/objects.py +0 -0
  43. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/parser.py +0 -0
  44. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/__init__.py +0 -0
  45. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/conformance.py +0 -0
  46. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/evaluate.py +0 -0
  47. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/flatten.py +0 -0
  48. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/model.py +0 -0
  49. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/__init__.py +0 -0
  50. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/catboost.py +0 -0
  51. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/json_dump.py +0 -0
  52. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/lightgbm.py +0 -0
  53. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/sklearn.py +0 -0
  54. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/xgboost.py +0 -0
  55. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/mining.py +0 -0
  56. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/objective.py +0 -0
  57. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/plausibility.py +0 -0
  58. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/py.typed +0 -0
  59. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/regions.py +0 -0
  60. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/targets.py +0 -0
  61. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/viz.py +0 -0
  62. {treecf-0.2.1 → treecf-0.2.2}/src/treecf/viz_batch.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: treecf
3
- Version: 0.2.1
3
+ Version: 0.2.2
4
4
  Classifier: Development Status :: 4 - Beta
5
5
  Classifier: Intended Audience :: Science/Research
6
6
  Classifier: License :: OSI Approved :: MIT License
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "treecf"
3
- version = "0.2.1"
3
+ version = "0.2.2"
4
4
  description = "Constrained, threshold-aware counterfactual explanations for tree ensembles (XGBoost, LightGBM, CatBoost, sklearn) — fast Rust genetic search, exact optimality proofs, certified infeasibility, and recourse regions."
5
5
  readme = "README.md"
6
6
  license = { text = "MIT" }
@@ -419,7 +419,7 @@ checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca"
419
419
 
420
420
  [[package]]
421
421
  name = "treecf-core"
422
- version = "0.2.1"
422
+ version = "0.2.2"
423
423
  dependencies = [
424
424
  "numpy",
425
425
  "pyo3",
@@ -1,6 +1,6 @@
1
1
  [package]
2
2
  name = "treecf-core"
3
- version = "0.2.1"
3
+ version = "0.2.2"
4
4
  edition = "2021"
5
5
  # f64::next_down (cells.rs) stabilized in 1.86; pyo3 0.29 needs 1.83
6
6
  rust-version = "1.86"
@@ -10,6 +10,7 @@ from treecf._errors import (
10
10
  UnsupportedModelError,
11
11
  )
12
12
  from treecf.api import Counterfactual, Explainer, Grid, Infeasible
13
+ from treecf.audit import constraints_fingerprint, ir_fingerprint
13
14
  from treecf.batch import BatchRecord, BatchResult
14
15
  from treecf.constraints import (
15
16
  AllowMissing,
@@ -27,7 +28,7 @@ from treecf.plausibility import Plausibility
27
28
  from treecf.regions import RecourseRegion
28
29
  from treecf.targets import Target
29
30
 
30
- __version__ = "0.2.1"
31
+ __version__ = "0.2.2"
31
32
 
32
33
  __all__ = [
33
34
  "AllowMissing",
@@ -58,5 +59,7 @@ __all__ = [
58
59
  "UnsupportedModelError",
59
60
  "__version__",
60
61
  "constraint",
62
+ "constraints_fingerprint",
63
+ "ir_fingerprint",
61
64
  "suggest_constraints",
62
65
  ]
@@ -1248,6 +1248,98 @@ class Explainer:
1248
1248
  )
1249
1249
  return self._region_for(x, x_cf, interval)
1250
1250
 
1251
+ def certificate(
1252
+ self,
1253
+ x: FloatArray,
1254
+ result: Counterfactual | Infeasible,
1255
+ target: Target,
1256
+ *,
1257
+ band: str | None = None,
1258
+ seed: int | None = None,
1259
+ node_budget: int | None = None,
1260
+ gap: float | None = None,
1261
+ time_budget_s: float | None = None,
1262
+ warm_start: bool | None = None,
1263
+ ) -> dict[str, object]:
1264
+ """Issue an audit certificate for a stored result (post-hoc, like
1265
+ ``recourse_region``).
1266
+
1267
+ A certificate is a reproducibility record plus a fresh verification —
1268
+ it binds the result to a model fingerprint, a constraint fingerprint,
1269
+ and the solve parameters, and re-verifies the returned plan at issue
1270
+ time; it does not cryptographically prove that a search ran or that a
1271
+ ``proof="optimal"`` claim is true — re-running with the recorded seed
1272
+ and budgets on a fingerprint-matching model is how a validator checks
1273
+ that. See
1274
+ [Certification — audit certificates](concepts/certification.md#audit-certificates)
1275
+ for the schema.
1276
+
1277
+ The certificate is a plain ``dict`` (``"schema_version": 1``) that
1278
+ serializes with ``json.dumps(cert, allow_nan=False, sort_keys=True)``;
1279
+ non-finite floats are encoded as the strings ``"NaN"``/``"Infinity"``/
1280
+ ``"-Infinity"``. Accepts a ``Counterfactual`` or an ``Infeasible`` —
1281
+ the certified "no" is exactly the case a validator cares most about.
1282
+ The verification block is computed fresh here, never copied from the
1283
+ solve: the plan's score, target membership, and constraint check are
1284
+ recomputed (plus the plausibility bound when configured, and a sampled
1285
+ set of region points when the result carries a region). A certificate
1286
+ whose fresh verification fails is still returned, with the failing
1287
+ booleans recorded — but a ``TreecfWarning`` names the failed check.
1288
+
1289
+ ``seed``/``node_budget``/``gap``/``time_budget_s``/``warm_start`` are
1290
+ recorded under ``solve.declared`` when given: the result object does
1291
+ not carry them, so they are caller-supplied, and the block's name
1292
+ makes that provenance explicit.
1293
+
1294
+ Args:
1295
+ x: The factual instance the result was solved from.
1296
+ result: The ``Counterfactual`` or ``Infeasible`` to certify.
1297
+ target: The target the result was solved against.
1298
+ band: For a ``Target.bands`` result, the band this result belongs
1299
+ to; required then, invalid otherwise.
1300
+ seed: The seed the solve ran with, if the caller wants it recorded.
1301
+ node_budget: The node budget the solve ran with, likewise.
1302
+ gap: The relative gap the solve ran with, likewise.
1303
+ time_budget_s: The time budget the solve ran with, likewise.
1304
+ warm_start: The warm-start setting the solve ran with, likewise.
1305
+
1306
+ Returns:
1307
+ The certificate as a strict-JSON-serializable ``dict``.
1308
+
1309
+ Raises:
1310
+ TreecfError: If ``target`` is a ``Target.bands`` ladder and
1311
+ ``band`` is missing or unknown, or if ``band`` is given for a
1312
+ plain-interval target.
1313
+ """
1314
+ from treecf.audit import build_certificate
1315
+
1316
+ return build_certificate(
1317
+ self, x, result, target, band=band, seed=seed, node_budget=node_budget,
1318
+ gap=gap, time_budget_s=time_budget_s, warm_start=warm_start,
1319
+ )
1320
+
1321
+ def check_certificate(self, cert: dict[str, object]) -> dict[str, object]:
1322
+ """Validate a stored certificate against *this* explainer.
1323
+
1324
+ Recomputes both fingerprints (model and constraints) against this
1325
+ explainer and re-runs the certificate's verification block from its
1326
+ stored factual/plan, so a tampered ``x_cf``, a swapped model, or a
1327
+ changed constraint set each flips the corresponding boolean. This
1328
+ method reports — it never raises on a mismatch.
1329
+
1330
+ Args:
1331
+ cert: A certificate produced by ``Explainer.certificate`` (a
1332
+ ``json.loads`` round trip of one works identically).
1333
+
1334
+ Returns:
1335
+ ``{"model_match": bool, "constraints_match": bool,
1336
+ "verification_ok": bool, "mismatches": [...]}`` with one
1337
+ human-readable string per mismatch.
1338
+ """
1339
+ from treecf.audit import check_certificate
1340
+
1341
+ return check_certificate(self, cert)
1342
+
1251
1343
  def _region_for(
1252
1344
  self, x: FloatArray, x_cf: FloatArray, interval: tuple[float, float]
1253
1345
  ) -> RecourseRegion:
@@ -0,0 +1,525 @@
1
+ """Audit certificates: reproducibility records with a fresh verification.
2
+
3
+ A certificate binds a returned result to a model fingerprint, a constraint
4
+ fingerprint, and the solve parameters, and re-verifies the returned plan at
5
+ issue time — it does not cryptographically prove that a search ran or that an
6
+ optimality claim is true; re-running with the recorded seed and budgets
7
+ against a fingerprint-matching model is how a validator checks that.
8
+
9
+ Certificates are plain ``dict[str, object]`` values (``"schema_version": 1``)
10
+ that serialize with ``json.dumps(cert, allow_nan=False, sort_keys=True)``:
11
+ non-finite floats are encoded as the strings ``"NaN"``, ``"Infinity"``, and
12
+ ``"-Infinity"`` wherever they can occur. See
13
+ [Certification — audit certificates](concepts/certification.md#audit-certificates).
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import hashlib
19
+ import math
20
+ import struct
21
+ import warnings
22
+ from datetime import UTC, datetime
23
+ from typing import TYPE_CHECKING
24
+
25
+ import numpy as np
26
+ import numpy.typing as npt
27
+
28
+ from treecf._errors import TreecfError, TreecfWarning
29
+ from treecf.constraints.objects import (
30
+ AllowMissing,
31
+ Constraint,
32
+ Equals,
33
+ Freeze,
34
+ Implies,
35
+ Linear,
36
+ Monotone,
37
+ OneHot,
38
+ Range,
39
+ )
40
+ from treecf.ir.evaluate import raw_score
41
+ from treecf.ir.model import EnsembleIR, SplitOp
42
+
43
+ if TYPE_CHECKING:
44
+ from treecf.api import Counterfactual, Explainer, Infeasible
45
+ from treecf.targets import Target
46
+
47
+ FloatArray = npt.NDArray[np.float64]
48
+
49
+ _NONE_U32 = 0xFFFFFFFF # sentinel for an index field that does not apply
50
+ _NONE_F64 = b"\xff" * 8 # sentinel for a float field that does not apply
51
+
52
+
53
+ def _json_float(value: float) -> float | str:
54
+ """A float as strict JSON allows it: finite as-is, non-finite as a string."""
55
+ if math.isnan(value):
56
+ return "NaN"
57
+ if value == math.inf:
58
+ return "Infinity"
59
+ if value == -math.inf:
60
+ return "-Infinity"
61
+ return float(value)
62
+
63
+
64
+ def _from_json_float(value: object) -> float:
65
+ """Inverse of ``_json_float``."""
66
+ if value == "NaN":
67
+ return math.nan
68
+ if value == "Infinity":
69
+ return math.inf
70
+ if value == "-Infinity":
71
+ return -math.inf
72
+ if isinstance(value, bool) or not isinstance(value, int | float):
73
+ raise TreecfError(f"not a certificate float: {value!r}")
74
+ return float(value)
75
+
76
+
77
+ def _encode_array(values: FloatArray) -> list[float | str]:
78
+ return [_json_float(v) for v in values.tolist()]
79
+
80
+
81
+ def _decode_array(values: object) -> FloatArray:
82
+ if not isinstance(values, list):
83
+ raise TreecfError("certificate array field is not a list")
84
+ return np.array([_from_json_float(v) for v in values], dtype=np.float64)
85
+
86
+
87
+ def ir_fingerprint(ir: EnsembleIR) -> str:
88
+ """SHA-256 fingerprint of an ensemble over a canonical byte encoding.
89
+
90
+ The encoding is positional bytes — the link name, the base score, the
91
+ tree count, then every node of every tree in index order with fixed-width
92
+ little-endian fields and fixed sentinel bytes where a field does not
93
+ apply to the node kind — so the fingerprint is stable across Python
94
+ versions, platforms, and dict ordering, and changes when any structural
95
+ or numeric detail of the ensemble changes (a one-ulp leaf perturbation
96
+ included).
97
+
98
+ Args:
99
+ ir: The parsed ensemble to fingerprint (``Explainer.ir``).
100
+
101
+ Returns:
102
+ A 64-character SHA-256 hex digest.
103
+ """
104
+ hasher = hashlib.sha256()
105
+ hasher.update(ir.link.name.encode("utf-8") + b"\x00")
106
+ hasher.update(struct.pack("<d", ir.base_score))
107
+ hasher.update(struct.pack("<I", len(ir.trees)))
108
+ for tree in ir.trees:
109
+ hasher.update(struct.pack("<I", len(tree.nodes)))
110
+ for node in tree.nodes:
111
+ if node.feature is None: # leaf
112
+ assert node.value is not None
113
+ hasher.update(b"\x00" + struct.pack("<I", _NONE_U32) + _NONE_F64)
114
+ hasher.update(b"\x00\x02") # op / missing_left sentinels
115
+ hasher.update(struct.pack("<II", _NONE_U32, _NONE_U32))
116
+ hasher.update(struct.pack("<d", node.value))
117
+ else:
118
+ assert node.threshold is not None and node.op is not None
119
+ assert node.missing_left is not None
120
+ assert node.left is not None and node.right is not None
121
+ hasher.update(b"\x01" + struct.pack("<I", node.feature))
122
+ hasher.update(struct.pack("<d", node.threshold))
123
+ hasher.update(
124
+ bytes((1 if node.op is SplitOp.LT else 2, 1 if node.missing_left else 0))
125
+ )
126
+ hasher.update(struct.pack("<II", node.left, node.right))
127
+ hasher.update(_NONE_F64)
128
+ return hasher.hexdigest()
129
+
130
+
131
+ def _constraint_record(c: Constraint, index: dict[str, int]) -> bytes:
132
+ """One constraint as canonical bytes: type tag, resolved feature indices,
133
+ parameters as little-endian f64. Never ``repr`` — bytes stay stable."""
134
+ if isinstance(c, Freeze):
135
+ return b"Freeze\x00" + struct.pack("<I", index[c.feature])
136
+ if isinstance(c, Monotone):
137
+ return b"Monotone\x00" + struct.pack("<I", index[c.feature]) + c.direction.encode("utf-8")
138
+ if isinstance(c, Range):
139
+ return b"Range\x00" + struct.pack("<Idd", index[c.feature], c.lo, c.hi)
140
+ if isinstance(c, Linear):
141
+ terms = sorted((index[name], float(coef)) for name, coef in c.coefficients.items())
142
+ body = b"".join(struct.pack("<Id", j, coef) for j, coef in terms)
143
+ return (
144
+ b"Linear\x00" + body + c.op.encode("utf-8") + b"\x00"
145
+ + struct.pack("<d", c.rhs) + c.missing_policy.encode("utf-8")
146
+ )
147
+ if isinstance(c, Equals):
148
+ return b"Equals\x00" + struct.pack("<Id", index[c.feature], c.value)
149
+ if isinstance(c, Implies):
150
+ return (
151
+ b"Implies\x00"
152
+ + struct.pack("<Id", index[c.condition.feature], c.condition.value)
153
+ + struct.pack("<Id", index[c.consequence.feature], c.consequence.value)
154
+ )
155
+ if isinstance(c, OneHot):
156
+ members = sorted(index[name] for name in c.features)
157
+ return b"OneHot\x00" + b"".join(struct.pack("<I", j) for j in members)
158
+ assert isinstance(c, AllowMissing)
159
+ delta_from = c.delta_miss if c.delta_from_miss is None else c.delta_from_miss
160
+ return b"AllowMissing\x00" + struct.pack("<Idd", index[c.feature], c.delta_miss, delta_from)
161
+
162
+
163
+ def _constraints_encoding(explainer: Explainer) -> tuple[bytes, str | None]:
164
+ """Canonical bytes for the effective objective and constraint set, plus a
165
+ non-reproducibility reason when a component has no canonical encoding."""
166
+ index = {name: j for j, name in enumerate(explainer.ir.feature_names)}
167
+ records = sorted(_constraint_record(c, index) for c in explainer.compiled.constraints)
168
+ parts = [struct.pack("<I", len(records)), *records]
169
+ parts.append(np.asarray(explainer.sigma, dtype="<f8").tobytes())
170
+ parts.append(np.asarray(explainer.weights, dtype="<f8").tobytes())
171
+ reason: str | None = None
172
+ for name in sorted(explainer.value_policy):
173
+ policy = explainer.value_policy[name]
174
+ parts.append(struct.pack("<I", index[name]))
175
+ if isinstance(policy, str):
176
+ parts.append(policy.encode("utf-8") + b"\x00")
177
+ elif callable(policy):
178
+ # a callable has no canonical encoding — never hash its repr
179
+ parts.append(b"unhashable_custom\x00")
180
+ reason = (
181
+ f"value_policy[{name!r}] is a callable policy with no canonical "
182
+ "encoding, so the constraints fingerprint cannot pin it"
183
+ )
184
+ else:
185
+ parts.append(b"Grid\x00" + struct.pack("<dd", policy.step, policy.anchor))
186
+ return b"".join(parts), reason
187
+
188
+
189
+ def constraints_fingerprint(explainer: Explainer) -> str:
190
+ """SHA-256 fingerprint of an explainer's effective objective and constraints.
191
+
192
+ Covers the compiled constraint set (type tags, resolved feature indices,
193
+ parameters), the distance normalizers ``sigma``, the per-feature
194
+ ``weights``, and every value-policy entry, all as canonical little-endian
195
+ bytes. A callable value policy has no canonical encoding: it is hashed as
196
+ a fixed ``unhashable_custom`` tag, and any certificate built from the
197
+ explainer records ``"reproducible": false`` with a reason.
198
+
199
+ Args:
200
+ explainer: The explainer whose constraint set to fingerprint.
201
+
202
+ Returns:
203
+ A 64-character SHA-256 hex digest.
204
+ """
205
+ encoding, _ = _constraints_encoding(explainer)
206
+ return hashlib.sha256(encoding).hexdigest()
207
+
208
+
209
+ def _backend_of(stats: dict[str, object]) -> str:
210
+ """Recover which backend family produced a result from its own stats;
211
+ ``"unknown"`` when the stats identify neither family."""
212
+ if "completed" in stats: # the exact backend always reports this key
213
+ return "exact"
214
+ if not stats or "generations" in stats:
215
+ return "genetic/python"
216
+ return "unknown"
217
+
218
+
219
+ def _json_stats(stats: dict[str, object]) -> dict[str, object]:
220
+ """A JSON-safe copy of solver stats (numpy scalars unwrapped, floats
221
+ encoded per the non-finite rule, anything exotic stringified)."""
222
+ out: dict[str, object] = {}
223
+ for key, value in stats.items():
224
+ if isinstance(value, np.generic):
225
+ value = value.item()
226
+ if isinstance(value, bool | int | str):
227
+ out[key] = value
228
+ elif isinstance(value, float):
229
+ out[key] = _json_float(value)
230
+ else:
231
+ out[key] = repr(value)
232
+ return out
233
+
234
+
235
+ def _region_points(
236
+ x_cf: FloatArray,
237
+ intervals: dict[str, tuple[float, float]],
238
+ feature_names: tuple[str, ...],
239
+ ) -> list[tuple[str, FloatArray]]:
240
+ """The finite point set a certificate verifies for a region (all corners
241
+ would be exponential): each widened feature's own two endpoints holding
242
+ every other coordinate at ``x_cf``, plus the all-lo and all-hi corners."""
243
+ index = {name: j for j, name in enumerate(feature_names)}
244
+ points: list[tuple[str, FloatArray]] = []
245
+ all_lo = x_cf.copy()
246
+ all_hi = x_cf.copy()
247
+ for name, (lo, hi) in intervals.items():
248
+ j = index[name]
249
+ for side, endpoint in (("lo", lo), ("hi", hi)):
250
+ point = x_cf.copy()
251
+ point[j] = endpoint
252
+ points.append((f"{name}={side}", point))
253
+ all_lo[j] = lo
254
+ all_hi[j] = hi
255
+ points.append(("all-lo", all_lo))
256
+ points.append(("all-hi", all_hi))
257
+ return points
258
+
259
+
260
+ def _verify_plan(
261
+ explainer: Explainer,
262
+ x: FloatArray,
263
+ x_cf: FloatArray,
264
+ interval: tuple[float, float],
265
+ region_intervals: dict[str, tuple[float, float]] | None,
266
+ ) -> tuple[dict[str, object], list[str]]:
267
+ """Fresh verification of a plan: recomputed score, target membership, the
268
+ compiled constraint check, plausibility when configured, and the region's
269
+ sampled points. Returns the block and the names of any failed checks."""
270
+ score = raw_score(explainer.ir, x_cf)
271
+ in_target = bool(interval[0] <= score <= interval[1])
272
+ constraints_ok = bool(explainer.compiled.check_matrix(x_cf[None, :], x)[0])
273
+ verification: dict[str, object] = {
274
+ "score_raw": _json_float(score),
275
+ "in_target_interval": in_target,
276
+ "constraints_ok": constraints_ok,
277
+ }
278
+ failed = [
279
+ name
280
+ for name, ok in (("in_target_interval", in_target), ("constraints_ok", constraints_ok))
281
+ if not ok
282
+ ]
283
+ if explainer.plausibility is not None:
284
+ anomaly = explainer.plausibility.anomaly_score(x_cf)
285
+ plausibility_ok = bool(anomaly <= explainer.plausibility.max_anomaly_score + 1e-12)
286
+ verification["plausibility_ok"] = plausibility_ok
287
+ if not plausibility_ok:
288
+ failed.append("plausibility_ok")
289
+ if region_intervals is not None:
290
+ checked: list[dict[str, object]] = []
291
+ for label, point in _region_points(x_cf, region_intervals, explainer.ir.feature_names):
292
+ ok = explainer._verify(x, point, interval) is None
293
+ checked.append({"point": label, "ok": ok})
294
+ if not ok:
295
+ failed.append(f"region point {label}")
296
+ verification["region_points"] = checked
297
+ return verification, failed
298
+
299
+
300
+ def _verify_infeasible(
301
+ explainer: Explainer, x: FloatArray, interval: tuple[float, float]
302
+ ) -> dict[str, object]:
303
+ """For an ``Infeasible`` there is no plan to verify — record only whether
304
+ the factual itself already sits outside the target interval."""
305
+ score = raw_score(explainer.ir, x)
306
+ return {
307
+ "factual_score_raw": _json_float(score),
308
+ "factual_in_target_interval": bool(interval[0] <= score <= interval[1]),
309
+ }
310
+
311
+
312
+ def _target_block(
313
+ target: Target, band: str | None, explainer: Explainer
314
+ ) -> tuple[dict[str, object], tuple[float, float]]:
315
+ """The certificate's target block and the resolved raw interval."""
316
+ if target.bands_spec is not None:
317
+ if band is None:
318
+ raise TreecfError(
319
+ "certificate for a Target.bands result requires band= naming which "
320
+ "band the result belongs to"
321
+ )
322
+ by_name = {name: (lo, hi) for name, lo, hi in target.bands_spec}
323
+ if band not in by_name:
324
+ raise TreecfError(f"band {band!r} is not part of this Target.bands ladder")
325
+ lo, hi = by_name[band]
326
+ interval = target.band_intervals(explainer.ir.link)[band]
327
+ else:
328
+ if band is not None:
329
+ raise TreecfError("band= is only valid for a Target.bands ladder")
330
+ lo, hi = target.lo, target.hi
331
+ interval = target.raw_interval(explainer.ir.link)
332
+ block: dict[str, object] = {
333
+ "space": target.space,
334
+ "lo": _json_float(lo),
335
+ "hi": _json_float(hi),
336
+ "raw_interval": [_json_float(interval[0]), _json_float(interval[1])],
337
+ }
338
+ if band is not None:
339
+ block["band"] = band
340
+ if target.space == "calibrated":
341
+ block["calibrator"] = "external — not embedded"
342
+ return block, interval
343
+
344
+
345
+ def build_certificate(
346
+ explainer: Explainer,
347
+ x: FloatArray,
348
+ result: Counterfactual | Infeasible,
349
+ target: Target,
350
+ *,
351
+ band: str | None = None,
352
+ seed: int | None = None,
353
+ node_budget: int | None = None,
354
+ gap: float | None = None,
355
+ time_budget_s: float | None = None,
356
+ warm_start: bool | None = None,
357
+ ) -> dict[str, object]:
358
+ """Body of ``Explainer.certificate``; see its docstring."""
359
+ from treecf import __version__
360
+ from treecf.api import Counterfactual
361
+
362
+ x = np.asarray(x, dtype=np.float64)
363
+ target_block, interval = _target_block(target, band, explainer)
364
+ _, reproducible_reason = _constraints_encoding(explainer)
365
+
366
+ model: dict[str, object] = {
367
+ "ir_fingerprint": ir_fingerprint(explainer.ir),
368
+ "feature_names": list(explainer.ir.feature_names),
369
+ "link": explainer.ir.link.name,
370
+ }
371
+ if explainer.plausibility is not None:
372
+ model["plausibility"] = {
373
+ "ir_fingerprint": ir_fingerprint(explainer.plausibility.if_ir),
374
+ "min_total_path": _json_float(explainer.plausibility.min_total_path),
375
+ }
376
+
377
+ declared: dict[str, object] = {}
378
+ if seed is not None:
379
+ declared["seed"] = seed
380
+ if node_budget is not None:
381
+ declared["node_budget"] = node_budget
382
+ if gap is not None:
383
+ declared["gap"] = _json_float(gap)
384
+ if time_budget_s is not None:
385
+ declared["time_budget_s"] = _json_float(time_budget_s)
386
+ if warm_start is not None:
387
+ declared["warm_start"] = warm_start
388
+ solve: dict[str, object] = {
389
+ "backend": _backend_of(result.solver_stats),
390
+ "proof": result.proof,
391
+ "solver_stats": _json_stats(result.solver_stats),
392
+ }
393
+ if declared:
394
+ solve["declared"] = declared
395
+
396
+ cert: dict[str, object] = {
397
+ "schema_version": 1,
398
+ "created_utc": datetime.now(UTC).isoformat(timespec="seconds"),
399
+ "treecf_version": __version__,
400
+ "reproducible": reproducible_reason is None,
401
+ "model": model,
402
+ "constraints": {
403
+ "fingerprint": constraints_fingerprint(explainer),
404
+ "listing": [repr(c) for c in explainer.compiled.constraints],
405
+ },
406
+ "target": target_block,
407
+ "solve": solve,
408
+ "factual": {"x": _encode_array(x)},
409
+ }
410
+ if reproducible_reason is not None:
411
+ cert["reproducible_reason"] = reproducible_reason
412
+
413
+ if isinstance(result, Counterfactual):
414
+ plan: dict[str, object] = {
415
+ "x_cf": _encode_array(result.x_cf),
416
+ "changes": {
417
+ name: [_json_float(src), _json_float(dst)]
418
+ for name, (src, dst) in result.changes.items()
419
+ },
420
+ "distance": _json_float(result.distance),
421
+ "snapped": dict(result.snapped),
422
+ }
423
+ region_intervals = None
424
+ if result.region is not None:
425
+ region_intervals = dict(result.region.feature_intervals)
426
+ plan["region_feature_intervals"] = {
427
+ name: [_json_float(lo), _json_float(hi)]
428
+ for name, (lo, hi) in region_intervals.items()
429
+ }
430
+ cert["plan"] = plan
431
+ verification, failed = _verify_plan(explainer, x, result.x_cf, interval, region_intervals)
432
+ if failed:
433
+ warnings.warn(
434
+ "certificate verification failed: " + ", ".join(failed)
435
+ + "; the certificate is issued with the failing checks recorded",
436
+ TreecfWarning,
437
+ stacklevel=3, # build_certificate <- Explainer.certificate <- user code
438
+ )
439
+ else:
440
+ cert["infeasible"] = {"reason": result.reason, "proof": result.proof}
441
+ verification = _verify_infeasible(explainer, x, interval)
442
+ cert["verification"] = verification
443
+ return cert
444
+
445
+
446
+ def check_certificate(explainer: Explainer, cert: dict[str, object]) -> dict[str, object]:
447
+ """Body of ``Explainer.check_certificate``; see its docstring."""
448
+ mismatches: list[str] = []
449
+
450
+ model = cert.get("model")
451
+ model_match = isinstance(model, dict) and model.get("ir_fingerprint") == ir_fingerprint(
452
+ explainer.ir
453
+ )
454
+ if not model_match:
455
+ mismatches.append("model fingerprint does not match this explainer's ensemble")
456
+ stored_plaus = model.get("plausibility") if isinstance(model, dict) else None
457
+ if explainer.plausibility is None:
458
+ if stored_plaus is not None:
459
+ model_match = False
460
+ mismatches.append(
461
+ "certificate declares a plausibility ensemble; this explainer has none"
462
+ )
463
+ elif not (
464
+ isinstance(stored_plaus, dict)
465
+ and stored_plaus.get("ir_fingerprint") == ir_fingerprint(explainer.plausibility.if_ir)
466
+ ):
467
+ model_match = False
468
+ mismatches.append("plausibility ensemble fingerprint does not match this explainer's")
469
+
470
+ constraints = cert.get("constraints")
471
+ constraints_match = (
472
+ isinstance(constraints, dict)
473
+ and constraints.get("fingerprint") == constraints_fingerprint(explainer)
474
+ )
475
+ if not constraints_match:
476
+ mismatches.append("constraints fingerprint does not match this explainer's constraint set")
477
+
478
+ verification_ok = True
479
+ try:
480
+ target = cert.get("target")
481
+ if not isinstance(target, dict):
482
+ raise TreecfError("certificate has no target block")
483
+ raw = target.get("raw_interval")
484
+ if not isinstance(raw, list) or len(raw) != 2:
485
+ raise TreecfError("certificate target block has no raw_interval")
486
+ interval = (_from_json_float(raw[0]), _from_json_float(raw[1]))
487
+ factual = cert.get("factual")
488
+ if not isinstance(factual, dict):
489
+ raise TreecfError("certificate has no factual block")
490
+ x = _decode_array(factual.get("x"))
491
+ plan = cert.get("plan")
492
+ if isinstance(plan, dict):
493
+ x_cf = _decode_array(plan.get("x_cf"))
494
+ intervals: dict[str, tuple[float, float]] | None = None
495
+ stored_intervals = plan.get("region_feature_intervals")
496
+ if isinstance(stored_intervals, dict):
497
+ intervals = {
498
+ str(name): (_from_json_float(pair[0]), _from_json_float(pair[1]))
499
+ for name, pair in stored_intervals.items()
500
+ }
501
+ _, failed = _verify_plan(explainer, x, x_cf, interval, intervals)
502
+ if failed:
503
+ verification_ok = False
504
+ mismatches.extend(f"verification failed: {name}" for name in failed)
505
+ else:
506
+ fresh = _verify_infeasible(explainer, x, interval)
507
+ stored = cert.get("verification")
508
+ stored_flag = stored.get("factual_in_target_interval") if isinstance(
509
+ stored, dict
510
+ ) else None
511
+ if stored_flag != fresh["factual_in_target_interval"]:
512
+ verification_ok = False
513
+ mismatches.append(
514
+ "the recomputed factual-in-target check no longer matches the certificate"
515
+ )
516
+ except Exception as exc: # a validator's tool reports; it never raises on bad input
517
+ verification_ok = False
518
+ mismatches.append(f"verification could not be re-run: {exc}")
519
+
520
+ return {
521
+ "model_match": model_match,
522
+ "constraints_match": constraints_match,
523
+ "verification_ok": verification_ok,
524
+ "mismatches": mismatches,
525
+ }
@@ -39,9 +39,10 @@ class BatchRecord:
39
39
  """One counterfactual (or the infeasibility marker) for one dataset row.
40
40
 
41
41
  Fields mirror ``Counterfactual`` (``x_cf``, ``changes``, ``distance``,
42
- ``n_changed``, ``score_raw``, ``score_prob``, ``region``), plus batch
43
- bookkeeping: ``id`` and ``k`` place the record in the dataset, and
44
- ``feasible`` distinguishes a real plan from the infeasibility marker.
42
+ ``n_changed``, ``score_raw``, ``score_prob``, ``proof``, ``solver_stats``,
43
+ ``region``), plus batch bookkeeping: ``id`` and ``k`` place the record in
44
+ the dataset, and ``feasible`` distinguishes a real plan from the
45
+ infeasibility marker.
45
46
 
46
47
  Attributes:
47
48
  id: The row identifier this record belongs to (an element of
@@ -84,6 +85,14 @@ class BatchRecord:
84
85
  region: The certified box around ``x_cf``, set only when
85
86
  ``explain_batch`` ran with ``region=True`` and ``feasible`` is
86
87
  ``True``; ``None`` otherwise.
88
+ proof: The claim this record makes, mirroring the single-instance
89
+ result that produced it: ``Counterfactual.proof`` (``"heuristic"``
90
+ | ``"optimal"`` | ``"optimal_within_gap"``) for a feasible
91
+ record, ``Infeasible.proof`` (``"search_exhausted"`` |
92
+ ``"certified"``) for an infeasibility marker.
93
+ solver_stats: Exact-backend diagnostics for the solve behind this
94
+ record, same keys as ``Counterfactual.solver_stats``; empty for
95
+ genetic/python solves (those engines report no per-row stats).
87
96
  """
88
97
 
89
98
  id: object
@@ -99,6 +108,8 @@ class BatchRecord:
99
108
  blocked_lever: str | None = None # diversity="lever-blocking": the frozen lever
100
109
  coalition: str | None = None # diversity="coalitions": the group this plan may touch
101
110
  region: RecourseRegion | None = None # set by explain_batch(..., region=True)
111
+ proof: str = "heuristic" # mirrors Counterfactual.proof / Infeasible.proof
112
+ solver_stats: dict[str, object] = field(default_factory=dict) # exact-backend only
102
113
 
103
114
 
104
115
  @dataclass(frozen=True)
@@ -180,6 +191,11 @@ class BatchResult:
180
191
  "seed": record.seed,
181
192
  "blocked_lever": record.blocked_lever,
182
193
  "coalition": record.coalition,
194
+ "proof": record.proof,
195
+ "solver_stats": {
196
+ key: encode_floats(value)
197
+ for key, value in record.solver_stats.items()
198
+ },
183
199
  "region": (
184
200
  None
185
201
  if record.region is None
@@ -207,7 +223,9 @@ class BatchResult:
207
223
  A file saved without ``region=True``, or by a version of treecf
208
224
  before regions existed, loads with every record's ``region`` set to
209
225
  ``None``; a file saved before coalition support loads with every
210
- record's ``coalition`` set to ``None``.
226
+ record's ``coalition`` set to ``None``; a file saved before per-record
227
+ proofs existed loads with ``proof`` defaulted by feasibility
228
+ (``"heuristic"``/``"search_exhausted"``) and empty ``solver_stats``.
211
229
 
212
230
  Args:
213
231
  path: Path to a file written by ``save``.
@@ -257,6 +275,14 @@ class BatchResult:
257
275
  blocked_lever=raw["blocked_lever"],
258
276
  coalition=raw.get("coalition"), # absent in pre-coalition files
259
277
  region=region,
278
+ # pre-0.2.2 files carry neither field; default by feasibility
279
+ proof=raw.get(
280
+ "proof", "heuristic" if raw["feasible"] else "search_exhausted"
281
+ ),
282
+ solver_stats={
283
+ key: decode_floats(value)
284
+ for key, value in raw.get("solver_stats", {}).items()
285
+ },
260
286
  )
261
287
  )
262
288
  essential_ids = [decode_floats(k) for k in data.get("essential_lever_ids", [])]
@@ -271,11 +297,13 @@ class BatchResult:
271
297
  def to_frame(self) -> Any:
272
298
  """One row per (id, k), wide ``cf_<feature>`` columns (pandas, lazy import).
273
299
 
274
- Every ``BatchRecord`` field except ``x_cf``/``changes``/``region``
275
- becomes its own column; ``x_cf`` is spread into one ``cf_<feature>``
276
- column per model feature (``NaN`` for an infeasible record, or an
277
- unchanged feature's factual-equal value); ``changes`` is summarized as
278
- a ``changed_features`` column (sorted feature names).
300
+ Every ``BatchRecord`` field except ``x_cf``/``changes``/``region``/
301
+ ``solver_stats`` becomes its own column (``solver_stats`` stays
302
+ record-only — read it off the ``BatchRecord`` directly); ``x_cf`` is
303
+ spread into one ``cf_<feature>`` column per model feature (``NaN`` for
304
+ an infeasible record, or an unchanged feature's factual-equal value);
305
+ ``changes`` is summarized as a ``changed_features`` column (sorted
306
+ feature names).
279
307
 
280
308
  Returns:
281
309
  A pandas ``DataFrame`` with one row per record.
@@ -300,6 +328,7 @@ class BatchResult:
300
328
  "seed": record.seed,
301
329
  "blocked_lever": record.blocked_lever,
302
330
  "coalition": record.coalition,
331
+ "proof": record.proof,
303
332
  "changed_features": sorted(record.changes),
304
333
  }
305
334
  for j, name in enumerate(self.feature_names):
@@ -540,6 +569,13 @@ def explain_batch(
540
569
  )
541
570
 
542
571
 
572
+ def _exact_stats(stats: dict[str, object]) -> dict[str, object]:
573
+ """Stats worth mirroring onto a record: the exact backend's per-solve
574
+ diagnostics (recognized by their ``completed`` key). Genetic/python engine
575
+ stats are not mirrored — those engines report no per-row diagnostics."""
576
+ return stats if "completed" in stats else {}
577
+
578
+
543
579
  def _record_from(
544
580
  row_id: object,
545
581
  k: int,
@@ -563,14 +599,23 @@ def _record_from(
563
599
  blocked_lever=blocked_lever,
564
600
  coalition=coalition,
565
601
  region=region,
602
+ proof=cf.proof,
603
+ solver_stats=_exact_stats(cf.solver_stats),
566
604
  )
567
605
 
568
606
 
569
- def _infeasible_record(row_id: object, k: int = 0, coalition: str | None = None) -> BatchRecord:
607
+ def _infeasible_record(
608
+ row_id: object,
609
+ k: int = 0,
610
+ coalition: str | None = None,
611
+ infeasible: Infeasible | None = None,
612
+ ) -> BatchRecord:
570
613
  return BatchRecord(
571
614
  id=row_id, k=k, feasible=False, x_cf=None, changes={},
572
615
  distance=None, n_changed=None, score_raw=None, score_prob=None,
573
616
  coalition=coalition,
617
+ proof="search_exhausted" if infeasible is None else infeasible.proof,
618
+ solver_stats={} if infeasible is None else _exact_stats(infeasible.solver_stats),
574
619
  )
575
620
 
576
621
 
@@ -860,8 +905,9 @@ def _rows_by_coalitions(
860
905
  records.append(_record_from(row_id, k, cf, coalition=name, region=reg))
861
906
  k += 1
862
907
  for name in solvers:
863
- if not isinstance(outcomes[name][i], Counterfactual):
864
- records.append(_infeasible_record(row_id, k=k, coalition=name))
908
+ outcome = outcomes[name][i]
909
+ if isinstance(outcome, Infeasible):
910
+ records.append(_infeasible_record(row_id, k=k, coalition=name, infeasible=outcome))
865
911
  k += 1
866
912
  return records
867
913
 
@@ -888,9 +934,10 @@ def _row_by_seeds(
888
934
  ``warm_start=False`` alongside it, so a ``None`` incumbent -- an
889
935
  infeasible warm draw -- also runs unwarmed rather than falling back to a
890
936
  per-attempt genetic pass; see ``explain_batch``'s ``allow_exact_batch``)."""
891
- from treecf.api import Counterfactual
937
+ from treecf.api import Counterfactual, Infeasible
892
938
 
893
939
  found: dict[frozenset[str], tuple[Counterfactual, int]] = {}
940
+ last_infeasible: Infeasible | None = None
894
941
  for attempt in range(_SEED_ATTEMPT_FACTOR * n_per_example):
895
942
  attempt_seed = master_seed + attempt
896
943
  result = explainer._explain(
@@ -905,8 +952,10 @@ def _row_by_seeds(
905
952
  found[key] = (result, attempt_seed)
906
953
  if len(found) == n_per_example:
907
954
  break
955
+ elif isinstance(result, Infeasible): # bands are rejected by explain_batch
956
+ last_infeasible = result
908
957
  if not found:
909
- return [_infeasible_record(row_id)]
958
+ return [_infeasible_record(row_id, infeasible=last_infeasible)]
910
959
  ranked = sorted(found.values(), key=lambda pair: pair[0].distance)[:n_per_example]
911
960
  interval = target.raw_interval(explainer.ir.link) if region else None
912
961
  return [
@@ -947,7 +996,7 @@ def _row_by_lever_blocking(
947
996
  assert not isinstance(explained, dict) # bands are rejected by explain_batch
948
997
  primary = explained
949
998
  if not isinstance(primary, Counterfactual):
950
- return [_infeasible_record(row_id)], []
999
+ return [_infeasible_record(row_id, infeasible=primary)], []
951
1000
 
952
1001
  interval = target.raw_interval(explainer.ir.link) if region else None
953
1002
  primary_region = (
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes