treecf 0.2.3__tar.gz → 0.3.0__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.3 → treecf-0.3.0}/PKG-INFO +14 -1
- {treecf-0.2.3 → treecf-0.3.0}/README.md +8 -0
- {treecf-0.2.3 → treecf-0.3.0}/pyproject.toml +9 -2
- {treecf-0.2.3 → treecf-0.3.0}/rust/Cargo.lock +1 -1
- {treecf-0.2.3 → treecf-0.3.0}/rust/Cargo.toml +1 -1
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/cells.rs +179 -1
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/constraints.rs +41 -0
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/exact/domains.rs +179 -14
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/exact/search.rs +176 -35
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/exact/test_support.rs +12 -1
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/ga.rs +92 -18
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/ir.rs +162 -0
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/lib.rs +3 -1
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/py.rs +114 -7
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/regions.rs +329 -26
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/__init__.py +5 -1
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/_errors.py +5 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/aim/cells.py +52 -1
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/api.py +328 -192
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/audit.py +206 -20
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/_exact_bounds.py +6 -1
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/_exact_domains.py +115 -6
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/exact.py +138 -49
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/exact_rust.py +4 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/genetic.py +43 -10
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/genetic_rust.py +13 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/regions_rust.py +30 -3
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/batch.py +182 -83
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/constraints/__init__.py +2 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/constraints/compile.py +105 -4
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/constraints/flatten.py +14 -0
- treecf-0.3.0/src/treecf/constraints/objects.py +203 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/constraints/parser.py +25 -19
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/conformance.py +19 -3
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/evaluate.py +58 -8
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/flatten.py +54 -5
- treecf-0.3.0/src/treecf/ir/model.py +157 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/parsers/__init__.py +17 -8
- treecf-0.3.0/src/treecf/ir/parsers/_catboost_cat.py +179 -0
- treecf-0.3.0/src/treecf/ir/parsers/catboost.py +380 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/parsers/json_dump.py +9 -5
- treecf-0.3.0/src/treecf/ir/parsers/lightgbm.py +212 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/parsers/sklearn.py +166 -17
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/parsers/xgboost.py +68 -11
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/mining.py +102 -75
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/objective.py +12 -2
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/plausibility.py +36 -24
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/regions.py +263 -50
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/targets.py +158 -110
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/viz.py +527 -113
- treecf-0.3.0/src/treecf/viz_batch.py +623 -0
- treecf-0.2.3/src/treecf/constraints/objects.py +0 -142
- treecf-0.2.3/src/treecf/ir/model.py +0 -59
- treecf-0.2.3/src/treecf/ir/parsers/catboost.py +0 -127
- treecf-0.2.3/src/treecf/ir/parsers/lightgbm.py +0 -121
- treecf-0.2.3/src/treecf/viz_batch.py +0 -338
- {treecf-0.2.3 → treecf-0.3.0}/LICENSE +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/exact/mod.rs +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/exact/orderpairs.rs +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/exact/propagation.rs +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/rust/src/interrupt.rs +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/_json.py +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/aim/__init__.py +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/__init__.py +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/_exact_orderpairs.py +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/backends/_exact_propagation.py +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/ir/__init__.py +0 -0
- {treecf-0.2.3 → treecf-0.3.0}/src/treecf/py.typed +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: treecf
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.3.0
|
|
4
4
|
Classifier: Development Status :: 4 - Beta
|
|
5
5
|
Classifier: Intended Audience :: Science/Research
|
|
6
6
|
Classifier: License :: OSI Approved :: MIT License
|
|
@@ -29,6 +29,7 @@ Requires-Dist: lightgbm>=4.3 ; extra == 'dev'
|
|
|
29
29
|
Requires-Dist: catboost>=1.2 ; extra == 'dev'
|
|
30
30
|
Requires-Dist: scikit-learn>=1.4 ; extra == 'dev'
|
|
31
31
|
Requires-Dist: matplotlib>=3.8 ; extra == 'dev'
|
|
32
|
+
Requires-Dist: probcal>=0.2 ; extra == 'dev'
|
|
32
33
|
Requires-Dist: mkdocs>=1.6 ; extra == 'docs'
|
|
33
34
|
Requires-Dist: mkdocs-material>=9.5 ; extra == 'docs'
|
|
34
35
|
Requires-Dist: mkdocstrings[python]>=0.27 ; extra == 'docs'
|
|
@@ -37,6 +38,9 @@ Requires-Dist: mkdocs-jupyter>=0.24 ; extra == 'docs'
|
|
|
37
38
|
Requires-Dist: ipykernel>=6.29 ; extra == 'docs'
|
|
38
39
|
Requires-Dist: lightgbm>=4.3 ; extra == 'lightgbm'
|
|
39
40
|
Requires-Dist: scikit-learn>=1.4 ; extra == 'sklearn'
|
|
41
|
+
Requires-Dist: pytest>=8.0 ; extra == 'test'
|
|
42
|
+
Requires-Dist: hypothesis>=6.100 ; extra == 'test'
|
|
43
|
+
Requires-Dist: probcal>=0.2 ; extra == 'test'
|
|
40
44
|
Requires-Dist: matplotlib>=3.8 ; extra == 'viz'
|
|
41
45
|
Requires-Dist: xgboost>=2.0 ; extra == 'xgboost'
|
|
42
46
|
Provides-Extra: all
|
|
@@ -45,6 +49,7 @@ Provides-Extra: dev
|
|
|
45
49
|
Provides-Extra: docs
|
|
46
50
|
Provides-Extra: lightgbm
|
|
47
51
|
Provides-Extra: sklearn
|
|
52
|
+
Provides-Extra: test
|
|
48
53
|
Provides-Extra: viz
|
|
49
54
|
Provides-Extra: xgboost
|
|
50
55
|
License-File: LICENSE
|
|
@@ -61,6 +66,8 @@ Project-URL: Issues, https://github.com/wlazlod/treecf/issues
|
|
|
61
66
|
|
|
62
67
|
# treecf
|
|
63
68
|
|
|
69
|
+
[](https://doi.org/10.5281/zenodo.22069503)
|
|
70
|
+
|
|
64
71
|
**Constrained, threshold-aware counterfactual explanations for tree ensembles.**
|
|
65
72
|
|
|
66
73
|
`treecf` answers the question: *"what is the minimal, feasible change to this instance such
|
|
@@ -124,6 +131,12 @@ proved = exp.explain(x, target=t, backend="exact") # proof="optimal", a cer
|
|
|
124
131
|
boxed = exp.explain(x, target=t, region=True) # res.region.describe() -> "utilization <= 0.4"
|
|
125
132
|
```
|
|
126
133
|
|
|
134
|
+
## Contributing
|
|
135
|
+
|
|
136
|
+
See [CONTRIBUTING.md](CONTRIBUTING.md) for dev setup, the test layers, and the
|
|
137
|
+
project's hard invariants; report security issues privately per
|
|
138
|
+
[SECURITY.md](SECURITY.md).
|
|
139
|
+
|
|
127
140
|
## License
|
|
128
141
|
|
|
129
142
|
MIT
|
|
@@ -1,5 +1,7 @@
|
|
|
1
1
|
# treecf
|
|
2
2
|
|
|
3
|
+
[](https://doi.org/10.5281/zenodo.22069503)
|
|
4
|
+
|
|
3
5
|
**Constrained, threshold-aware counterfactual explanations for tree ensembles.**
|
|
4
6
|
|
|
5
7
|
`treecf` answers the question: *"what is the minimal, feasible change to this instance such
|
|
@@ -63,6 +65,12 @@ proved = exp.explain(x, target=t, backend="exact") # proof="optimal", a cer
|
|
|
63
65
|
boxed = exp.explain(x, target=t, region=True) # res.region.describe() -> "utilization <= 0.4"
|
|
64
66
|
```
|
|
65
67
|
|
|
68
|
+
## Contributing
|
|
69
|
+
|
|
70
|
+
See [CONTRIBUTING.md](CONTRIBUTING.md) for dev setup, the test layers, and the
|
|
71
|
+
project's hard invariants; report security issues privately per
|
|
72
|
+
[SECURITY.md](SECURITY.md).
|
|
73
|
+
|
|
66
74
|
## License
|
|
67
75
|
|
|
68
76
|
MIT
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "treecf"
|
|
3
|
-
version = "0.
|
|
3
|
+
version = "0.3.0"
|
|
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" }
|
|
@@ -44,6 +44,11 @@ all = [
|
|
|
44
44
|
"scikit-learn>=1.4",
|
|
45
45
|
"matplotlib>=3.8",
|
|
46
46
|
]
|
|
47
|
+
test = [
|
|
48
|
+
"pytest>=8.0",
|
|
49
|
+
"hypothesis>=6.100",
|
|
50
|
+
"probcal>=0.2",
|
|
51
|
+
]
|
|
47
52
|
dev = [
|
|
48
53
|
"maturin>=1.7",
|
|
49
54
|
"pytest>=8.0",
|
|
@@ -56,6 +61,7 @@ dev = [
|
|
|
56
61
|
"catboost>=1.2",
|
|
57
62
|
"scikit-learn>=1.4",
|
|
58
63
|
"matplotlib>=3.8",
|
|
64
|
+
"probcal>=0.2",
|
|
59
65
|
]
|
|
60
66
|
docs = [
|
|
61
67
|
"mkdocs>=1.6",
|
|
@@ -115,10 +121,11 @@ ignore_missing_imports = true
|
|
|
115
121
|
|
|
116
122
|
[tool.pytest.ini_options]
|
|
117
123
|
testpaths = ["tests"]
|
|
118
|
-
addopts = "-q --strict-markers -m 'not bench'"
|
|
124
|
+
addopts = "-q --strict-markers -m 'not bench and not slow'"
|
|
119
125
|
markers = [
|
|
120
126
|
"bench: non-gating performance smoke benchmarks (run with -m bench)",
|
|
121
127
|
"rust: cross-language tests requiring the dev Rust extension (maturin develop)",
|
|
128
|
+
"slow: docs snippet/structure harnesses (run with -m slow)",
|
|
122
129
|
]
|
|
123
130
|
filterwarnings = [
|
|
124
131
|
"error",
|
|
@@ -152,7 +152,9 @@ pub fn feature_cells_joint(ensembles: &[&Ensemble]) -> Vec<Vec<Cell>> {
|
|
|
152
152
|
let mut pairs: Vec<Vec<(f64, bool)>> = vec![Vec::new(); n_features];
|
|
153
153
|
for ens in ensembles {
|
|
154
154
|
for i in 0..ens.feature.len() {
|
|
155
|
-
|
|
155
|
+
// set-membership splits carry no threshold; their features partition
|
|
156
|
+
// into category blocks instead of interval cells
|
|
157
|
+
if ens.feature[i] >= 0 && ens.node_set[i] < 0 {
|
|
156
158
|
pairs[ens.feature[i] as usize].push((ens.threshold[i], ens.is_lt[i]));
|
|
157
159
|
}
|
|
158
160
|
}
|
|
@@ -160,6 +162,78 @@ pub fn feature_cells_joint(ensembles: &[&Ensemble]) -> Vec<Vec<Cell>> {
|
|
|
160
162
|
pairs.iter().map(|p| build_cells(p)).collect()
|
|
161
163
|
}
|
|
162
164
|
|
|
165
|
+
/// Category blocks per feature across ensembles — port of
|
|
166
|
+
/// `treecf.aim.cells.category_blocks`. Indexed by feature; numeric features get
|
|
167
|
+
/// an empty vec. Two codes share a block iff no split in any ensemble separates
|
|
168
|
+
/// them: set-membership splits partition by membership, numeric splits on a
|
|
169
|
+
/// categorical feature (an isolation forest trained on raw codes) partition by
|
|
170
|
+
/// threshold side. Blocks are ordered by smallest member (codes scan
|
|
171
|
+
/// ascending), members ascend, and a block's representative is its first entry.
|
|
172
|
+
///
|
|
173
|
+
/// Panics where Python raises `ValueError` (cardinality disagreement), matching
|
|
174
|
+
/// the `feature_cells_joint` convention.
|
|
175
|
+
pub fn category_blocks_joint(ensembles: &[&Ensemble]) -> Vec<Vec<Vec<u32>>> {
|
|
176
|
+
let n_features = ensembles[0].n_features;
|
|
177
|
+
let mut cardinality = vec![0u32; n_features];
|
|
178
|
+
for ens in ensembles {
|
|
179
|
+
for (j, slot) in cardinality.iter_mut().enumerate() {
|
|
180
|
+
let k = ens.cardinality[j];
|
|
181
|
+
if k > 0 {
|
|
182
|
+
if *slot == 0 {
|
|
183
|
+
*slot = k;
|
|
184
|
+
} else {
|
|
185
|
+
assert_eq!(
|
|
186
|
+
*slot, k,
|
|
187
|
+
"ensembles disagree on the cardinality of feature {j}"
|
|
188
|
+
);
|
|
189
|
+
}
|
|
190
|
+
}
|
|
191
|
+
}
|
|
192
|
+
}
|
|
193
|
+
let mut out = Vec::with_capacity(n_features);
|
|
194
|
+
for (j, &card) in cardinality.iter().enumerate() {
|
|
195
|
+
let k = card as usize;
|
|
196
|
+
if k == 0 {
|
|
197
|
+
out.push(Vec::new());
|
|
198
|
+
continue;
|
|
199
|
+
}
|
|
200
|
+
let mut signatures: Vec<Vec<bool>> = vec![Vec::new(); k];
|
|
201
|
+
for ens in ensembles {
|
|
202
|
+
for i in 0..ens.feature.len() {
|
|
203
|
+
if ens.feature[i] != j as i32 {
|
|
204
|
+
continue;
|
|
205
|
+
}
|
|
206
|
+
if ens.node_set[i] >= 0 {
|
|
207
|
+
for (code, sig) in signatures.iter_mut().enumerate() {
|
|
208
|
+
sig.push(ens.set_contains(ens.node_set[i], code as f64));
|
|
209
|
+
}
|
|
210
|
+
} else if ens.is_lt[i] {
|
|
211
|
+
for (code, sig) in signatures.iter_mut().enumerate() {
|
|
212
|
+
sig.push((code as f64) < ens.threshold[i]);
|
|
213
|
+
}
|
|
214
|
+
} else {
|
|
215
|
+
for (code, sig) in signatures.iter_mut().enumerate() {
|
|
216
|
+
sig.push((code as f64) <= ens.threshold[i]);
|
|
217
|
+
}
|
|
218
|
+
}
|
|
219
|
+
}
|
|
220
|
+
}
|
|
221
|
+
let mut blocks: Vec<Vec<u32>> = Vec::new();
|
|
222
|
+
let mut seen: Vec<usize> = Vec::new(); // block index -> a code carrying its signature
|
|
223
|
+
for code in 0..k {
|
|
224
|
+
match seen.iter().position(|&c| signatures[c] == signatures[code]) {
|
|
225
|
+
Some(b) => blocks[b].push(code as u32),
|
|
226
|
+
None => {
|
|
227
|
+
seen.push(code);
|
|
228
|
+
blocks.push(vec![code as u32]);
|
|
229
|
+
}
|
|
230
|
+
}
|
|
231
|
+
}
|
|
232
|
+
out.push(blocks);
|
|
233
|
+
}
|
|
234
|
+
out
|
|
235
|
+
}
|
|
236
|
+
|
|
163
237
|
/// Index of the unique cell containing `x` — port of `treecf.aim.cells.cell_index`.
|
|
164
238
|
///
|
|
165
239
|
/// Panics where Python raises `ValueError`. The cells partition the whole line,
|
|
@@ -175,6 +249,7 @@ pub fn cell_index(cells: &[Cell], x: f64) -> usize {
|
|
|
175
249
|
#[cfg(test)]
|
|
176
250
|
mod tests {
|
|
177
251
|
use super::*;
|
|
252
|
+
use crate::ir::Link;
|
|
178
253
|
|
|
179
254
|
#[test]
|
|
180
255
|
fn lt_le_collision_yields_singleton() {
|
|
@@ -252,4 +327,107 @@ mod tests {
|
|
|
252
327
|
assert_eq!(cell_index(&cells, 1.0), 1);
|
|
253
328
|
assert_eq!(cell_index(&cells, 1.5), 2);
|
|
254
329
|
}
|
|
330
|
+
|
|
331
|
+
/// One set stump per word list, all on feature 0, cardinality `k`.
|
|
332
|
+
fn set_ens(sets: &[u64], k: u32) -> Ensemble {
|
|
333
|
+
let n = sets.len().max(1);
|
|
334
|
+
let mut feature = Vec::new();
|
|
335
|
+
let mut node_set = Vec::new();
|
|
336
|
+
let mut left = Vec::new();
|
|
337
|
+
let mut right = Vec::new();
|
|
338
|
+
let mut value = Vec::new();
|
|
339
|
+
let mut roots = Vec::new();
|
|
340
|
+
for (t, _) in sets.iter().enumerate() {
|
|
341
|
+
let base = (t * 3) as u32;
|
|
342
|
+
roots.push(base);
|
|
343
|
+
feature.extend_from_slice(&[0, -1, -1]);
|
|
344
|
+
node_set.extend_from_slice(&[t as i32, -1, -1]);
|
|
345
|
+
left.extend_from_slice(&[base + 1, 0, 0]);
|
|
346
|
+
right.extend_from_slice(&[base + 2, 0, 0]);
|
|
347
|
+
value.extend_from_slice(&[0.0, -1.0, 1.0]);
|
|
348
|
+
}
|
|
349
|
+
if sets.is_empty() {
|
|
350
|
+
roots.push(0);
|
|
351
|
+
feature.push(-1);
|
|
352
|
+
node_set.push(-1);
|
|
353
|
+
left.push(0);
|
|
354
|
+
right.push(0);
|
|
355
|
+
value.push(0.0);
|
|
356
|
+
}
|
|
357
|
+
let n_nodes = feature.len();
|
|
358
|
+
Ensemble::new(
|
|
359
|
+
feature,
|
|
360
|
+
vec![0.0; n_nodes],
|
|
361
|
+
vec![false; n_nodes],
|
|
362
|
+
vec![false; n_nodes],
|
|
363
|
+
left,
|
|
364
|
+
right,
|
|
365
|
+
value,
|
|
366
|
+
roots,
|
|
367
|
+
0.0,
|
|
368
|
+
Link::Identity,
|
|
369
|
+
n,
|
|
370
|
+
)
|
|
371
|
+
.unwrap()
|
|
372
|
+
.with_categories(
|
|
373
|
+
node_set,
|
|
374
|
+
(0..=sets.len() as u32).collect(),
|
|
375
|
+
sets.to_vec(),
|
|
376
|
+
{
|
|
377
|
+
let mut c = vec![0; n];
|
|
378
|
+
c[0] = k;
|
|
379
|
+
c
|
|
380
|
+
},
|
|
381
|
+
)
|
|
382
|
+
.unwrap()
|
|
383
|
+
}
|
|
384
|
+
|
|
385
|
+
#[test]
|
|
386
|
+
fn no_splits_one_block() {
|
|
387
|
+
let e = set_ens(&[], 4);
|
|
388
|
+
assert_eq!(e.category_blocks()[0], vec![vec![0, 1, 2, 3]]);
|
|
389
|
+
}
|
|
390
|
+
|
|
391
|
+
#[test]
|
|
392
|
+
fn one_set_two_blocks_numbered_by_smallest_member() {
|
|
393
|
+
let e = set_ens(&[0b101], 4); // {0, 2}
|
|
394
|
+
assert_eq!(e.category_blocks()[0], vec![vec![0, 2], vec![1, 3]]);
|
|
395
|
+
}
|
|
396
|
+
|
|
397
|
+
#[test]
|
|
398
|
+
fn two_sets_refine_to_singletons() {
|
|
399
|
+
let e = set_ens(&[0b101, 0b1100], 4); // {0,2} then {2,3}
|
|
400
|
+
assert_eq!(
|
|
401
|
+
e.category_blocks()[0],
|
|
402
|
+
vec![vec![0], vec![1], vec![2], vec![3]]
|
|
403
|
+
);
|
|
404
|
+
}
|
|
405
|
+
|
|
406
|
+
#[test]
|
|
407
|
+
fn joint_blocks_refine_and_numeric_splits_partition_codes() {
|
|
408
|
+
let a = set_ens(&[0b101], 4);
|
|
409
|
+
// a numeric stump on the same feature at 1.5 (an isolation forest would)
|
|
410
|
+
let b = {
|
|
411
|
+
let mut e = Ensemble::new(
|
|
412
|
+
vec![0, -1, -1],
|
|
413
|
+
vec![1.5, 0.0, 0.0],
|
|
414
|
+
vec![false, false, false],
|
|
415
|
+
vec![false, false, false],
|
|
416
|
+
vec![1, 0, 0],
|
|
417
|
+
vec![2, 0, 0],
|
|
418
|
+
vec![0.0, -1.0, 1.0],
|
|
419
|
+
vec![0],
|
|
420
|
+
0.0,
|
|
421
|
+
Link::Identity,
|
|
422
|
+
1,
|
|
423
|
+
)
|
|
424
|
+
.unwrap();
|
|
425
|
+
e = e
|
|
426
|
+
.with_categories(vec![-1, -1, -1], vec![0], vec![], vec![4])
|
|
427
|
+
.unwrap();
|
|
428
|
+
e
|
|
429
|
+
};
|
|
430
|
+
let joint = category_blocks_joint(&[&a, &b]);
|
|
431
|
+
assert_eq!(joint[0], vec![vec![0], vec![1], vec![2], vec![3]]);
|
|
432
|
+
}
|
|
255
433
|
}
|
|
@@ -35,6 +35,27 @@ pub struct Constraints {
|
|
|
35
35
|
pub implications: Vec<(u32, f64, u32, f64)>, // cond_idx, cond_val, cons_idx, cons_val
|
|
36
36
|
pub onehot: Vec<Vec<u32>>,
|
|
37
37
|
pub allow_missing: Vec<(u32, f64, f64)>, // (feature, delta_to, delta_from), index-sorted
|
|
38
|
+
// (feature, allowed-code bitset words), index-sorted; an empty word list
|
|
39
|
+
// means the declared set is empty (nothing is allowed on that feature)
|
|
40
|
+
pub allowed_categories: Vec<(u32, Vec<u64>)>,
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
/// Membership of `v` in an allowed-code bitset: an integral code whose bit is
|
|
44
|
+
/// set. Non-integral values are never members; codes beyond the words are not.
|
|
45
|
+
#[inline]
|
|
46
|
+
pub(crate) fn code_allowed(words: &[u64], v: f64) -> bool {
|
|
47
|
+
if !v.is_finite() {
|
|
48
|
+
return false;
|
|
49
|
+
}
|
|
50
|
+
let code = v as i64;
|
|
51
|
+
if code as f64 != v || code < 0 {
|
|
52
|
+
return false;
|
|
53
|
+
}
|
|
54
|
+
let word = (code >> 6) as usize;
|
|
55
|
+
if word >= words.len() {
|
|
56
|
+
return false;
|
|
57
|
+
}
|
|
58
|
+
(words[word] >> (code & 63)) & 1 == 1
|
|
38
59
|
}
|
|
39
60
|
|
|
40
61
|
/// Python `max(a, b)`: returns b only if b > a. Two consequences both mirrors
|
|
@@ -184,6 +205,12 @@ impl Constraints {
|
|
|
184
205
|
return false; // NaN sum compares false -> infeasible, like numpy
|
|
185
206
|
}
|
|
186
207
|
}
|
|
208
|
+
for (j, words) in &self.allowed_categories {
|
|
209
|
+
let v = row[*j as usize];
|
|
210
|
+
if !v.is_nan() && !code_allowed(words, v) {
|
|
211
|
+
return false;
|
|
212
|
+
}
|
|
213
|
+
}
|
|
187
214
|
true
|
|
188
215
|
}
|
|
189
216
|
|
|
@@ -239,6 +266,19 @@ impl Constraints {
|
|
|
239
266
|
row[j] = v;
|
|
240
267
|
}
|
|
241
268
|
}
|
|
269
|
+
for (j, words) in &self.allowed_categories {
|
|
270
|
+
let j = *j as usize;
|
|
271
|
+
// keep an allowed code; else the smallest allowed code (deterministic)
|
|
272
|
+
if words.is_empty() {
|
|
273
|
+
continue; // nothing legal to write; check rejects
|
|
274
|
+
}
|
|
275
|
+
if !row[j].is_nan() && !code_allowed(words, row[j]) {
|
|
276
|
+
if let Some(word) = words.iter().position(|&w| w != 0) {
|
|
277
|
+
let bit = words[word].trailing_zeros() as usize;
|
|
278
|
+
row[j] = (word * 64 + bit) as f64;
|
|
279
|
+
}
|
|
280
|
+
}
|
|
281
|
+
}
|
|
242
282
|
// cyclic projection: boxes intersected with halfspaces, a few sweeps
|
|
243
283
|
// (must mirror repair_matrix float-for-float: same term order, one
|
|
244
284
|
// residual/denom division, lower-bound-first clipping)
|
|
@@ -345,6 +385,7 @@ mod tests {
|
|
|
345
385
|
implications: vec![],
|
|
346
386
|
onehot: vec![],
|
|
347
387
|
allow_missing: vec![],
|
|
388
|
+
allowed_categories: vec![],
|
|
348
389
|
}
|
|
349
390
|
}
|
|
350
391
|
|