treecf 0.2.4__tar.gz → 0.3.1__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.4 → treecf-0.3.1}/PKG-INFO +7 -1
- {treecf-0.2.4 → treecf-0.3.1}/README.md +6 -0
- {treecf-0.2.4 → treecf-0.3.1}/pyproject.toml +3 -2
- {treecf-0.2.4 → treecf-0.3.1}/rust/Cargo.lock +1 -1
- {treecf-0.2.4 → treecf-0.3.1}/rust/Cargo.toml +1 -1
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/cells.rs +179 -1
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/constraints.rs +41 -0
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/domains.rs +180 -15
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/mod.rs +6 -1
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/propagation.rs +32 -1
- treecf-0.3.1/rust/src/exact/refine.rs +825 -0
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/search.rs +589 -165
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/test_support.rs +12 -1
- treecf-0.3.1/rust/src/exact/trace.rs +108 -0
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/ga.rs +92 -18
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/ir.rs +162 -0
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/lib.rs +4 -1
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/py.rs +200 -16
- treecf-0.3.1/rust/src/region_slab.rs +344 -0
- treecf-0.3.1/rust/src/regions.rs +1551 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/__init__.py +8 -1
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/_errors.py +5 -0
- treecf-0.3.1/src/treecf/_menu.py +629 -0
- treecf-0.3.1/src/treecf/_portfolio.py +619 -0
- treecf-0.3.1/src/treecf/_region_slab.py +278 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/aim/cells.py +52 -1
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/api.py +690 -235
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/audit.py +134 -27
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_bounds.py +42 -3
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_domains.py +115 -6
- treecf-0.3.1/src/treecf/backends/_exact_profile.py +106 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_propagation.py +26 -1
- treecf-0.3.1/src/treecf/backends/_exact_refine.py +653 -0
- treecf-0.3.1/src/treecf/backends/_exact_trace.py +50 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/exact.py +346 -92
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/exact_rust.py +22 -1
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/genetic.py +43 -10
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/genetic_rust.py +13 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/regions_rust.py +65 -5
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/batch.py +279 -113
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/__init__.py +2 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/compile.py +105 -4
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/flatten.py +14 -0
- treecf-0.3.1/src/treecf/constraints/objects.py +203 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/parser.py +25 -19
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/conformance.py +19 -3
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/evaluate.py +58 -8
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/flatten.py +54 -5
- treecf-0.3.1/src/treecf/ir/model.py +248 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/parsers/__init__.py +17 -8
- treecf-0.3.1/src/treecf/ir/parsers/_catboost_cat.py +179 -0
- treecf-0.3.1/src/treecf/ir/parsers/_float32.py +49 -0
- treecf-0.3.1/src/treecf/ir/parsers/catboost.py +385 -0
- treecf-0.3.1/src/treecf/ir/parsers/json_dump.py +95 -0
- treecf-0.3.1/src/treecf/ir/parsers/lightgbm.py +214 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/parsers/sklearn.py +166 -43
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/parsers/xgboost.py +72 -12
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/mining.py +102 -75
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/objective.py +12 -2
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/plausibility.py +36 -24
- treecf-0.3.1/src/treecf/regions.py +1052 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/targets.py +158 -110
- treecf-0.3.1/src/treecf/viz.py +1635 -0
- treecf-0.3.1/src/treecf/viz_batch.py +623 -0
- treecf-0.2.4/rust/src/regions.rs +0 -757
- treecf-0.2.4/src/treecf/constraints/objects.py +0 -142
- treecf-0.2.4/src/treecf/ir/model.py +0 -59
- treecf-0.2.4/src/treecf/ir/parsers/catboost.py +0 -127
- treecf-0.2.4/src/treecf/ir/parsers/json_dump.py +0 -35
- treecf-0.2.4/src/treecf/ir/parsers/lightgbm.py +0 -121
- treecf-0.2.4/src/treecf/regions.py +0 -477
- treecf-0.2.4/src/treecf/viz.py +0 -868
- treecf-0.2.4/src/treecf/viz_batch.py +0 -338
- {treecf-0.2.4 → treecf-0.3.1}/LICENSE +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/orderpairs.rs +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/rust/src/interrupt.rs +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/_json.py +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/aim/__init__.py +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/__init__.py +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_orderpairs.py +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/__init__.py +0 -0
- {treecf-0.2.4 → treecf-0.3.1}/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.1
|
|
4
4
|
Classifier: Development Status :: 4 - Beta
|
|
5
5
|
Classifier: Intended Audience :: Science/Research
|
|
6
6
|
Classifier: License :: OSI Approved :: MIT License
|
|
@@ -131,6 +131,12 @@ proved = exp.explain(x, target=t, backend="exact") # proof="optimal", a cer
|
|
|
131
131
|
boxed = exp.explain(x, target=t, region=True) # res.region.describe() -> "utilization <= 0.4"
|
|
132
132
|
```
|
|
133
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
|
+
|
|
134
140
|
## License
|
|
135
141
|
|
|
136
142
|
MIT
|
|
@@ -65,6 +65,12 @@ proved = exp.explain(x, target=t, backend="exact") # proof="optimal", a cer
|
|
|
65
65
|
boxed = exp.explain(x, target=t, region=True) # res.region.describe() -> "utilization <= 0.4"
|
|
66
66
|
```
|
|
67
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
|
+
|
|
68
74
|
## License
|
|
69
75
|
|
|
70
76
|
MIT
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "treecf"
|
|
3
|
-
version = "0.
|
|
3
|
+
version = "0.3.1"
|
|
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" }
|
|
@@ -121,10 +121,11 @@ ignore_missing_imports = true
|
|
|
121
121
|
|
|
122
122
|
[tool.pytest.ini_options]
|
|
123
123
|
testpaths = ["tests"]
|
|
124
|
-
addopts = "-q --strict-markers -m 'not bench'"
|
|
124
|
+
addopts = "-q --strict-markers -m 'not bench and not slow'"
|
|
125
125
|
markers = [
|
|
126
126
|
"bench: non-gating performance smoke benchmarks (run with -m bench)",
|
|
127
127
|
"rust: cross-language tests requiring the dev Rust extension (maturin develop)",
|
|
128
|
+
"slow: docs snippet/structure harnesses (run with -m slow)",
|
|
128
129
|
]
|
|
129
130
|
filterwarnings = [
|
|
130
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
|
|
|
@@ -12,7 +12,7 @@ use crate::ir::Ensemble;
|
|
|
12
12
|
/// Python's `<`-based ordering: `-0.0` and `0.0` compare equal, so a stable sort
|
|
13
13
|
/// leaves them in insertion order. Costs and sort values are never NaN here.
|
|
14
14
|
#[inline]
|
|
15
|
-
fn py_cmp(a: f64, b: f64) -> std::cmp::Ordering {
|
|
15
|
+
pub(crate) fn py_cmp(a: f64, b: f64) -> std::cmp::Ordering {
|
|
16
16
|
use std::cmp::Ordering;
|
|
17
17
|
if a < b {
|
|
18
18
|
Ordering::Less
|
|
@@ -27,6 +27,9 @@ fn py_cmp(a: f64, b: f64) -> std::cmp::Ordering {
|
|
|
27
27
|
|
|
28
28
|
/// One feature's contribution to the objective — the per-feature term of
|
|
29
29
|
/// `genetic.objective()`, same four cases, same multiply-then-divide order.
|
|
30
|
+
/// A categorical change (`is_cat`) costs one flat unit in place of the
|
|
31
|
+
/// absolute code distance; NaN transitions keep their declared deltas.
|
|
32
|
+
#[allow(clippy::too_many_arguments)]
|
|
30
33
|
fn term_cost(
|
|
31
34
|
x_j: f64,
|
|
32
35
|
r: f64,
|
|
@@ -35,6 +38,7 @@ fn term_cost(
|
|
|
35
38
|
lam: f64,
|
|
36
39
|
to_miss: f64,
|
|
37
40
|
from_miss: f64,
|
|
41
|
+
is_cat: bool,
|
|
38
42
|
) -> f64 {
|
|
39
43
|
let x_nan = x_j.is_nan();
|
|
40
44
|
let r_nan = r.is_nan();
|
|
@@ -50,11 +54,12 @@ fn term_cost(
|
|
|
50
54
|
if r == x_j {
|
|
51
55
|
return 0.0;
|
|
52
56
|
}
|
|
53
|
-
let delta = (r - x_j).abs();
|
|
57
|
+
let delta = if is_cat { 1.0 } else { (r - x_j).abs() };
|
|
54
58
|
lam + (weight_j * delta) / sigma_j
|
|
55
59
|
}
|
|
56
60
|
|
|
57
61
|
/// Full-row objective, accumulated in ascending feature index.
|
|
62
|
+
/// `cardinality[j] > 0` marks feature j categorical (flat change cost).
|
|
58
63
|
pub(crate) fn cost_of_row(
|
|
59
64
|
x: &[f64],
|
|
60
65
|
row: &[f64],
|
|
@@ -62,11 +67,21 @@ pub(crate) fn cost_of_row(
|
|
|
62
67
|
weights: &[f64],
|
|
63
68
|
lam: f64,
|
|
64
69
|
deltas: &[(f64, f64)],
|
|
70
|
+
cardinality: &[u32],
|
|
65
71
|
) -> f64 {
|
|
66
72
|
let mut total = 0.0;
|
|
67
73
|
for j in 0..x.len() {
|
|
68
74
|
let (to_miss, from_miss) = deltas[j];
|
|
69
|
-
total += term_cost(
|
|
75
|
+
total += term_cost(
|
|
76
|
+
x[j],
|
|
77
|
+
row[j],
|
|
78
|
+
weights[j],
|
|
79
|
+
sigma[j],
|
|
80
|
+
lam,
|
|
81
|
+
to_miss,
|
|
82
|
+
from_miss,
|
|
83
|
+
cardinality[j] > 0,
|
|
84
|
+
);
|
|
70
85
|
}
|
|
71
86
|
total
|
|
72
87
|
}
|
|
@@ -376,6 +391,7 @@ pub(crate) fn build_domains(
|
|
|
376
391
|
weights: &[f64],
|
|
377
392
|
lam: f64,
|
|
378
393
|
policies: &[Option<ValuePolicy>],
|
|
394
|
+
blocks: &[Vec<Vec<u32>>],
|
|
379
395
|
) -> Vec<Vec<State>> {
|
|
380
396
|
let (lo, hi, frozen) = cons.instance_bounds(x);
|
|
381
397
|
let lo: Vec<f64> = lo
|
|
@@ -414,6 +430,28 @@ pub(crate) fn build_domains(
|
|
|
414
430
|
let (to_miss, from_miss) = deltas[j];
|
|
415
431
|
let (weight_j, sigma_j) = (weights[j], sigma[j]);
|
|
416
432
|
|
|
433
|
+
if !blocks[j].is_empty() {
|
|
434
|
+
let allowed = cons
|
|
435
|
+
.allowed_categories
|
|
436
|
+
.iter()
|
|
437
|
+
.find(|(idx, _)| *idx as usize == j)
|
|
438
|
+
.map(|(_, words)| words.as_slice());
|
|
439
|
+
domains.push(categorical_states(
|
|
440
|
+
x_j,
|
|
441
|
+
&blocks[j],
|
|
442
|
+
allowed,
|
|
443
|
+
frozen[j],
|
|
444
|
+
allow_j,
|
|
445
|
+
suppress_nan[j],
|
|
446
|
+
weight_j,
|
|
447
|
+
sigma_j,
|
|
448
|
+
lam,
|
|
449
|
+
to_miss,
|
|
450
|
+
from_miss,
|
|
451
|
+
));
|
|
452
|
+
continue;
|
|
453
|
+
}
|
|
454
|
+
|
|
417
455
|
if frozen[j] {
|
|
418
456
|
let idx = if x_nan {
|
|
419
457
|
cells.len()
|
|
@@ -429,11 +467,19 @@ pub(crate) fn build_domains(
|
|
|
429
467
|
if !x_nan {
|
|
430
468
|
// the pin fixes the only value the feature may take; going
|
|
431
469
|
// missing is a separate question AllowMissing still answers
|
|
432
|
-
let cost = term_cost(x_j, v, weight_j, sigma_j, lam, to_miss, from_miss);
|
|
470
|
+
let cost = term_cost(x_j, v, weight_j, sigma_j, lam, to_miss, from_miss, false);
|
|
433
471
|
let mut states = vec![State::new(v, cost, cell_index(cells, v), false)];
|
|
434
472
|
if allow_j && !suppress_nan[j] {
|
|
435
|
-
let nan_cost =
|
|
436
|
-
|
|
473
|
+
let nan_cost = term_cost(
|
|
474
|
+
x_j,
|
|
475
|
+
f64::NAN,
|
|
476
|
+
weight_j,
|
|
477
|
+
sigma_j,
|
|
478
|
+
lam,
|
|
479
|
+
to_miss,
|
|
480
|
+
from_miss,
|
|
481
|
+
false,
|
|
482
|
+
);
|
|
437
483
|
states.push(State::new(f64::NAN, nan_cost, cells.len(), true));
|
|
438
484
|
sort_states(&mut states);
|
|
439
485
|
}
|
|
@@ -449,7 +495,7 @@ pub(crate) fn build_domains(
|
|
|
449
495
|
states.push(State::new(f64::NAN, 0.0, cells.len(), true));
|
|
450
496
|
}
|
|
451
497
|
if allow_j {
|
|
452
|
-
let cost = term_cost(x_j, v, weight_j, sigma_j, lam, to_miss, from_miss);
|
|
498
|
+
let cost = term_cost(x_j, v, weight_j, sigma_j, lam, to_miss, from_miss, false);
|
|
453
499
|
states.push(State::new(v, cost, cell_index(cells, v), false));
|
|
454
500
|
}
|
|
455
501
|
sort_states(&mut states);
|
|
@@ -492,7 +538,8 @@ pub(crate) fn build_domains(
|
|
|
492
538
|
if keep_added && val == x_j {
|
|
493
539
|
continue;
|
|
494
540
|
}
|
|
495
|
-
let cost =
|
|
541
|
+
let cost =
|
|
542
|
+
term_cost(x_j, val, weight_j, sigma_j, lam, to_miss, from_miss, false);
|
|
496
543
|
states.push(State::new(val, cost, local_idx, false));
|
|
497
544
|
}
|
|
498
545
|
continue;
|
|
@@ -504,7 +551,7 @@ pub(crate) fn build_domains(
|
|
|
504
551
|
if !iv.contains(val) || (keep_added && val == x_j) {
|
|
505
552
|
continue;
|
|
506
553
|
}
|
|
507
|
-
let cost = term_cost(x_j, val, weight_j, sigma_j, lam, to_miss, from_miss);
|
|
554
|
+
let cost = term_cost(x_j, val, weight_j, sigma_j, lam, to_miss, from_miss, false);
|
|
508
555
|
states.push(State::new(val, cost, local_idx, false));
|
|
509
556
|
added_here.push(val);
|
|
510
557
|
}
|
|
@@ -523,14 +570,23 @@ pub(crate) fn build_domains(
|
|
|
523
570
|
if added_here.contains(&r) {
|
|
524
571
|
continue; // the demanded value was this cell's nearest point too
|
|
525
572
|
}
|
|
526
|
-
let cost = term_cost(x_j, r, weight_j, sigma_j, lam, to_miss, from_miss);
|
|
573
|
+
let cost = term_cost(x_j, r, weight_j, sigma_j, lam, to_miss, from_miss, false);
|
|
527
574
|
let mut state = State::new(r, cost, local_idx, false);
|
|
528
575
|
state.snapped = snapped;
|
|
529
576
|
states.push(state);
|
|
530
577
|
}
|
|
531
578
|
|
|
532
579
|
if allow_j && !suppress_nan[j] {
|
|
533
|
-
let nan_cost = term_cost(
|
|
580
|
+
let nan_cost = term_cost(
|
|
581
|
+
x_j,
|
|
582
|
+
f64::NAN,
|
|
583
|
+
weight_j,
|
|
584
|
+
sigma_j,
|
|
585
|
+
lam,
|
|
586
|
+
to_miss,
|
|
587
|
+
from_miss,
|
|
588
|
+
false,
|
|
589
|
+
);
|
|
534
590
|
states.push(State::new(f64::NAN, nan_cost, cells.len(), true));
|
|
535
591
|
}
|
|
536
592
|
|
|
@@ -575,15 +631,121 @@ fn referenced_features(cons: &Constraints) -> Vec<bool> {
|
|
|
575
631
|
for &(j, _, _) in &cons.allow_missing {
|
|
576
632
|
refs[j as usize] = true;
|
|
577
633
|
}
|
|
634
|
+
for (j, _) in &cons.allowed_categories {
|
|
635
|
+
refs[*j as usize] = true;
|
|
636
|
+
}
|
|
578
637
|
refs
|
|
579
638
|
}
|
|
580
639
|
|
|
640
|
+
/// Candidate states of a categorical feature — the mirror of
|
|
641
|
+
/// `_categorical_states`. A block's representative is its smallest member the
|
|
642
|
+
/// declared allowed set admits (cost is flat within a block, so the choice is
|
|
643
|
+
/// free; smallest is the determinism rule); a block with no admissible member
|
|
644
|
+
/// contributes no state, and an allowed set admitting nothing empties the
|
|
645
|
+
/// whole domain — the same certified-infeasible signal a contradictory
|
|
646
|
+
/// numeric pin produces. `cell_idx` carries the block index; the NaN state's
|
|
647
|
+
/// sentinel index is the block count.
|
|
648
|
+
#[allow(clippy::too_many_arguments)]
|
|
649
|
+
fn categorical_states(
|
|
650
|
+
x_j: f64,
|
|
651
|
+
blocks_j: &[Vec<u32>],
|
|
652
|
+
allowed: Option<&[u64]>,
|
|
653
|
+
frozen_j: bool,
|
|
654
|
+
allow_j: bool,
|
|
655
|
+
suppressed: bool,
|
|
656
|
+
weight_j: f64,
|
|
657
|
+
sigma_j: f64,
|
|
658
|
+
lam: f64,
|
|
659
|
+
to_miss: f64,
|
|
660
|
+
from_miss: f64,
|
|
661
|
+
) -> Vec<State> {
|
|
662
|
+
use crate::constraints::code_allowed;
|
|
663
|
+
|
|
664
|
+
let x_nan = x_j.is_nan();
|
|
665
|
+
let n_blocks = blocks_j.len();
|
|
666
|
+
let block_of = |code: f64| -> usize {
|
|
667
|
+
blocks_j
|
|
668
|
+
.iter()
|
|
669
|
+
.position(|block| block.contains(&(code as u32)))
|
|
670
|
+
.expect("code not covered by blocks (should be impossible)")
|
|
671
|
+
};
|
|
672
|
+
|
|
673
|
+
if frozen_j {
|
|
674
|
+
let idx = if x_nan { n_blocks } else { block_of(x_j) };
|
|
675
|
+
return vec![State::new(x_j, 0.0, idx, x_nan)];
|
|
676
|
+
}
|
|
677
|
+
|
|
678
|
+
if x_nan && !allow_j {
|
|
679
|
+
if suppressed {
|
|
680
|
+
return Vec::new();
|
|
681
|
+
}
|
|
682
|
+
return vec![State::new(x_j, 0.0, n_blocks, true)];
|
|
683
|
+
}
|
|
684
|
+
|
|
685
|
+
let is_allowed = |code: f64| -> bool {
|
|
686
|
+
match allowed {
|
|
687
|
+
None => true,
|
|
688
|
+
Some(words) => code_allowed(words, code),
|
|
689
|
+
}
|
|
690
|
+
};
|
|
691
|
+
|
|
692
|
+
let mut states: Vec<State> = Vec::new();
|
|
693
|
+
let mut keep_block = usize::MAX;
|
|
694
|
+
if !x_nan && is_allowed(x_j) {
|
|
695
|
+
keep_block = block_of(x_j);
|
|
696
|
+
states.push(State::new(x_j, 0.0, keep_block, false));
|
|
697
|
+
}
|
|
698
|
+
|
|
699
|
+
for (block_idx, block) in blocks_j.iter().enumerate() {
|
|
700
|
+
if block_idx == keep_block {
|
|
701
|
+
continue; // the unchanged factual already represents its block
|
|
702
|
+
}
|
|
703
|
+
let member = block.iter().copied().find(|&c| is_allowed(c as f64));
|
|
704
|
+
if let Some(code) = member {
|
|
705
|
+
let rep = code as f64;
|
|
706
|
+
let cost = term_cost(x_j, rep, weight_j, sigma_j, lam, to_miss, from_miss, true);
|
|
707
|
+
states.push(State::new(rep, cost, block_idx, false));
|
|
708
|
+
}
|
|
709
|
+
}
|
|
710
|
+
|
|
711
|
+
if allow_j && !suppressed {
|
|
712
|
+
let nan_cost = term_cost(
|
|
713
|
+
x_j,
|
|
714
|
+
f64::NAN,
|
|
715
|
+
weight_j,
|
|
716
|
+
sigma_j,
|
|
717
|
+
lam,
|
|
718
|
+
to_miss,
|
|
719
|
+
from_miss,
|
|
720
|
+
true,
|
|
721
|
+
);
|
|
722
|
+
states.push(State::new(f64::NAN, nan_cost, n_blocks, true));
|
|
723
|
+
}
|
|
724
|
+
|
|
725
|
+
sort_states(&mut states);
|
|
726
|
+
states
|
|
727
|
+
}
|
|
728
|
+
|
|
581
729
|
/// Search order: descending split count in the joint grid, ties ascending index.
|
|
582
730
|
/// A feature with no split anywhere and no referencing constraint is left out —
|
|
583
731
|
/// its domain is a single keep state, so it never needs to branch.
|
|
584
|
-
pub(crate) fn feature_order(
|
|
732
|
+
pub(crate) fn feature_order(
|
|
733
|
+
grids: &[Vec<Cell>],
|
|
734
|
+
cons: &Constraints,
|
|
735
|
+
blocks: &[Vec<Vec<u32>>],
|
|
736
|
+
) -> Vec<usize> {
|
|
585
737
|
let referenced = referenced_features(cons);
|
|
586
|
-
let split_counts: Vec<usize> = grids
|
|
738
|
+
let split_counts: Vec<usize> = grids
|
|
739
|
+
.iter()
|
|
740
|
+
.enumerate()
|
|
741
|
+
.map(|(j, cells)| {
|
|
742
|
+
if blocks[j].is_empty() {
|
|
743
|
+
cells.len() - 1
|
|
744
|
+
} else {
|
|
745
|
+
blocks[j].len() - 1
|
|
746
|
+
}
|
|
747
|
+
})
|
|
748
|
+
.collect();
|
|
587
749
|
let mut included: Vec<usize> = (0..grids.len())
|
|
588
750
|
.filter(|&j| split_counts[j] > 0 || referenced[j])
|
|
589
751
|
.collect();
|
|
@@ -1107,7 +1269,8 @@ mod tests {
|
|
|
1107
1269
|
let mut cons = cons_base(4);
|
|
1108
1270
|
cons.freeze = vec![3]; // referenced but split-free: kept, and last
|
|
1109
1271
|
let grids = constraint_cells(&cons, &[&ens]);
|
|
1110
|
-
|
|
1272
|
+
let blocks = crate::cells::category_blocks_joint(&[&ens]);
|
|
1273
|
+
assert_eq!(feature_order(&grids, &cons, &blocks), vec![1, 2, 3]);
|
|
1111
1274
|
}
|
|
1112
1275
|
|
|
1113
1276
|
#[test]
|
|
@@ -1115,6 +1278,7 @@ mod tests {
|
|
|
1115
1278
|
let ens = golden_ens();
|
|
1116
1279
|
let cons = cons_base(2);
|
|
1117
1280
|
let grids = constraint_cells(&cons, &[&ens]);
|
|
1281
|
+
let blocks = crate::cells::category_blocks_joint(&[&ens]);
|
|
1118
1282
|
let domains = build_domains(
|
|
1119
1283
|
&grids,
|
|
1120
1284
|
&[2.0, 0.0],
|
|
@@ -1123,8 +1287,9 @@ mod tests {
|
|
|
1123
1287
|
&[1.0, 1.0],
|
|
1124
1288
|
0.0,
|
|
1125
1289
|
&no_policies(2),
|
|
1290
|
+
&blocks,
|
|
1126
1291
|
);
|
|
1127
|
-
let order = feature_order(&grids, &cons);
|
|
1292
|
+
let order = feature_order(&grids, &cons, &blocks);
|
|
1128
1293
|
assert_eq!(h_suffix(&order, &domains), vec![0.0, 0.0]);
|
|
1129
1294
|
}
|
|
1130
1295
|
|