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.
- {treecf-0.2.1 → treecf-0.2.2}/PKG-INFO +1 -1
- {treecf-0.2.1 → treecf-0.2.2}/pyproject.toml +1 -1
- {treecf-0.2.1 → treecf-0.2.2}/rust/Cargo.lock +1 -1
- {treecf-0.2.1 → treecf-0.2.2}/rust/Cargo.toml +1 -1
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/__init__.py +4 -1
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/api.py +92 -0
- treecf-0.2.2/src/treecf/audit.py +525 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/batch.py +64 -15
- {treecf-0.2.1 → treecf-0.2.2}/LICENSE +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/README.md +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/cells.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/constraints.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/domains.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/mod.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/orderpairs.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/propagation.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/search.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/exact/test_support.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/ga.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/interrupt.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/ir.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/lib.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/py.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/rust/src/regions.rs +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/_errors.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/_json.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/aim/__init__.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/aim/cells.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/__init__.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_bounds.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_domains.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_orderpairs.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/_exact_propagation.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/exact.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/exact_rust.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/genetic.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/genetic_rust.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/backends/regions_rust.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/__init__.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/compile.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/flatten.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/objects.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/constraints/parser.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/__init__.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/conformance.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/evaluate.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/flatten.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/model.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/__init__.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/catboost.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/json_dump.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/lightgbm.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/sklearn.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/ir/parsers/xgboost.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/mining.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/objective.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/plausibility.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/py.typed +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/regions.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/targets.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/viz.py +0 -0
- {treecf-0.2.1 → treecf-0.2.2}/src/treecf/viz_batch.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "treecf"
|
|
3
|
-
version = "0.2.
|
|
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" }
|
|
@@ -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.
|
|
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``, ``
|
|
43
|
-
bookkeeping: ``id`` and ``k`` place the record in
|
|
44
|
-
``feasible`` distinguishes a real plan from the
|
|
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
|
|
276
|
-
|
|
277
|
-
|
|
278
|
-
|
|
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(
|
|
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
|
-
|
|
864
|
-
|
|
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
|
|
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
|