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.
Files changed (82) hide show
  1. {treecf-0.2.4 → treecf-0.3.1}/PKG-INFO +7 -1
  2. {treecf-0.2.4 → treecf-0.3.1}/README.md +6 -0
  3. {treecf-0.2.4 → treecf-0.3.1}/pyproject.toml +3 -2
  4. {treecf-0.2.4 → treecf-0.3.1}/rust/Cargo.lock +1 -1
  5. {treecf-0.2.4 → treecf-0.3.1}/rust/Cargo.toml +1 -1
  6. {treecf-0.2.4 → treecf-0.3.1}/rust/src/cells.rs +179 -1
  7. {treecf-0.2.4 → treecf-0.3.1}/rust/src/constraints.rs +41 -0
  8. {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/domains.rs +180 -15
  9. {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/mod.rs +6 -1
  10. {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/propagation.rs +32 -1
  11. treecf-0.3.1/rust/src/exact/refine.rs +825 -0
  12. {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/search.rs +589 -165
  13. {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/test_support.rs +12 -1
  14. treecf-0.3.1/rust/src/exact/trace.rs +108 -0
  15. {treecf-0.2.4 → treecf-0.3.1}/rust/src/ga.rs +92 -18
  16. {treecf-0.2.4 → treecf-0.3.1}/rust/src/ir.rs +162 -0
  17. {treecf-0.2.4 → treecf-0.3.1}/rust/src/lib.rs +4 -1
  18. {treecf-0.2.4 → treecf-0.3.1}/rust/src/py.rs +200 -16
  19. treecf-0.3.1/rust/src/region_slab.rs +344 -0
  20. treecf-0.3.1/rust/src/regions.rs +1551 -0
  21. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/__init__.py +8 -1
  22. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/_errors.py +5 -0
  23. treecf-0.3.1/src/treecf/_menu.py +629 -0
  24. treecf-0.3.1/src/treecf/_portfolio.py +619 -0
  25. treecf-0.3.1/src/treecf/_region_slab.py +278 -0
  26. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/aim/cells.py +52 -1
  27. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/api.py +690 -235
  28. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/audit.py +134 -27
  29. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_bounds.py +42 -3
  30. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_domains.py +115 -6
  31. treecf-0.3.1/src/treecf/backends/_exact_profile.py +106 -0
  32. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_propagation.py +26 -1
  33. treecf-0.3.1/src/treecf/backends/_exact_refine.py +653 -0
  34. treecf-0.3.1/src/treecf/backends/_exact_trace.py +50 -0
  35. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/exact.py +346 -92
  36. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/exact_rust.py +22 -1
  37. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/genetic.py +43 -10
  38. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/genetic_rust.py +13 -0
  39. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/regions_rust.py +65 -5
  40. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/batch.py +279 -113
  41. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/__init__.py +2 -0
  42. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/compile.py +105 -4
  43. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/flatten.py +14 -0
  44. treecf-0.3.1/src/treecf/constraints/objects.py +203 -0
  45. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/constraints/parser.py +25 -19
  46. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/conformance.py +19 -3
  47. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/evaluate.py +58 -8
  48. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/flatten.py +54 -5
  49. treecf-0.3.1/src/treecf/ir/model.py +248 -0
  50. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/parsers/__init__.py +17 -8
  51. treecf-0.3.1/src/treecf/ir/parsers/_catboost_cat.py +179 -0
  52. treecf-0.3.1/src/treecf/ir/parsers/_float32.py +49 -0
  53. treecf-0.3.1/src/treecf/ir/parsers/catboost.py +385 -0
  54. treecf-0.3.1/src/treecf/ir/parsers/json_dump.py +95 -0
  55. treecf-0.3.1/src/treecf/ir/parsers/lightgbm.py +214 -0
  56. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/parsers/sklearn.py +166 -43
  57. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/parsers/xgboost.py +72 -12
  58. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/mining.py +102 -75
  59. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/objective.py +12 -2
  60. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/plausibility.py +36 -24
  61. treecf-0.3.1/src/treecf/regions.py +1052 -0
  62. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/targets.py +158 -110
  63. treecf-0.3.1/src/treecf/viz.py +1635 -0
  64. treecf-0.3.1/src/treecf/viz_batch.py +623 -0
  65. treecf-0.2.4/rust/src/regions.rs +0 -757
  66. treecf-0.2.4/src/treecf/constraints/objects.py +0 -142
  67. treecf-0.2.4/src/treecf/ir/model.py +0 -59
  68. treecf-0.2.4/src/treecf/ir/parsers/catboost.py +0 -127
  69. treecf-0.2.4/src/treecf/ir/parsers/json_dump.py +0 -35
  70. treecf-0.2.4/src/treecf/ir/parsers/lightgbm.py +0 -121
  71. treecf-0.2.4/src/treecf/regions.py +0 -477
  72. treecf-0.2.4/src/treecf/viz.py +0 -868
  73. treecf-0.2.4/src/treecf/viz_batch.py +0 -338
  74. {treecf-0.2.4 → treecf-0.3.1}/LICENSE +0 -0
  75. {treecf-0.2.4 → treecf-0.3.1}/rust/src/exact/orderpairs.rs +0 -0
  76. {treecf-0.2.4 → treecf-0.3.1}/rust/src/interrupt.rs +0 -0
  77. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/_json.py +0 -0
  78. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/aim/__init__.py +0 -0
  79. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/__init__.py +0 -0
  80. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/backends/_exact_orderpairs.py +0 -0
  81. {treecf-0.2.4 → treecf-0.3.1}/src/treecf/ir/__init__.py +0 -0
  82. {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.2.4
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.2.4"
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",
@@ -419,7 +419,7 @@ checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca"
419
419
 
420
420
  [[package]]
421
421
  name = "treecf-core"
422
- version = "0.2.4"
422
+ version = "0.3.1"
423
423
  dependencies = [
424
424
  "numpy",
425
425
  "pyo3",
@@ -1,6 +1,6 @@
1
1
  [package]
2
2
  name = "treecf-core"
3
- version = "0.2.4"
3
+ version = "0.3.1"
4
4
  edition = "2021"
5
5
  # f64::next_down (cells.rs) stabilized in 1.86; pyo3 0.29 needs 1.83
6
6
  rust-version = "1.86"
@@ -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
- if ens.feature[i] >= 0 {
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(x[j], row[j], weights[j], sigma[j], lam, to_miss, from_miss);
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
- term_cost(x_j, f64::NAN, weight_j, sigma_j, lam, to_miss, from_miss);
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 = term_cost(x_j, val, weight_j, sigma_j, lam, to_miss, from_miss);
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(x_j, f64::NAN, weight_j, sigma_j, lam, to_miss, from_miss);
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(grids: &[Vec<Cell>], cons: &Constraints) -> Vec<usize> {
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.iter().map(|cells| cells.len() - 1).collect();
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
- assert_eq!(feature_order(&grids, &cons), vec![1, 2, 3]);
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