capability_grower_configuration_from_skeleton/
infer_level_skipping.rs1crate::ix!();
3
4pub trait TryInferLevelSkippingConfiguration {
5 fn try_infer_level_skipping_configuration(&self) -> Option<LevelSkippingConfiguration>;
6}
7
8impl TryInferLevelSkippingConfiguration for Skeleton {
9 fn try_infer_level_skipping_configuration(&self) -> Option<LevelSkippingConfiguration> {
10 let skipping_map = self.measure_tree_level_skipping();
12 if skipping_map.is_empty() {
13 return None;
14 }
15
16 let max_level = skipping_map.keys().copied().max().unwrap_or(0);
18
19 let mut skip_probs = Vec::with_capacity((max_level as usize) + 1);
21 for lvl in 0..=max_level {
22 let p = if let Some(stats) = skipping_map.get(&lvl) {
23 let nc = *stats.total_node_count();
24 let lc = *stats.total_leaf_count();
25 if nc == 0 {
26 0.0
27 } else {
28 (lc as f32 / nc as f32).powf(1.2)
29 }
30 } else {
31 0.0
32 };
33 skip_probs.push(p);
34 }
35
36 if max_level == u8::MAX {
39 if let Some(last) = skip_probs.last_mut() {
41 *last = 1.0;
42 }
43 }
44
45 LevelSkippingConfigurationBuilder::default()
47 .leaf_probability_per_level(skip_probs)
48 .build()
49 .ok()
50 }
51}
52
53#[cfg(test)]
54mod try_infer_level_skipping_configuration_tests {
55 use super::*;
56
57 #[traced_test]
59 fn empty_skeleton_returns_none() {
60 let skel = SkeletonBuilder::default().build().unwrap();
61 assert!(skel.try_infer_level_skipping_configuration().is_none());
62 }
63
64 #[traced_test]
66 fn single_leaf_returns_full_probability() {
67 let n0 = SkeletonNodeBuilder::default()
68 .id(0_u16)
69 .name("root")
70 .original_key("root")
71 .build(NodeKind::LeafHolder)
72 .unwrap();
73 let skel = SkeletonBuilder::default()
74 .nodes(vec![n0])
75 .root_id(Some(0_u16))
76 .build()
77 .unwrap();
78 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
79 let p = cfg.leaf_probability_per_level();
80 assert_eq!(p.len(), 1);
81 assert!((p[0] - 1.0).abs() < f32::EPSILON);
82 }
83
84 #[traced_test]
86 fn linear_chain_three_levels() {
87
88 let n0 = SkeletonNodeBuilder::default()
89 .id(0_u16)
90 .child_ids(vec![1_u16])
91 .name("n0")
92 .original_key("n0")
93 .build(NodeKind::Dispatch)
94 .unwrap();
95
96 let n1 = SkeletonNodeBuilder::default()
97 .id(1_u16)
98 .child_ids(vec![2_u16])
99 .name("n1")
100 .original_key("n1")
101 .build(NodeKind::Dispatch)
102 .unwrap();
103
104 let n2 = SkeletonNodeBuilder::default()
105 .id(2_u16)
106 .name("n2")
107 .original_key("n2")
108 .build(NodeKind::LeafHolder)
109 .unwrap();
110
111 let skel = SkeletonBuilder::default()
112 .nodes(vec![n0,n1,n2])
113 .root_id(Some(0_u16))
114 .build()
115 .unwrap();
116
117 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
118 let probs = cfg.leaf_probability_per_level();
119 assert_eq!(probs.len(), 3);
120 assert!((probs[0] - 0.0).abs() < 1e-6);
121 assert!((probs[1] - 0.0).abs() < 1e-6);
122 assert!((probs[2] - 1.0).abs() < 1e-6);
123 }
124
125 #[traced_test]
127 fn branch_with_two_leaves() {
128
129 let root = SkeletonNodeBuilder::default()
130 .id(0_u16)
131 .child_ids(vec![1_u16,2_u16])
132 .name("root")
133 .original_key("root")
134 .build(NodeKind::Dispatch)
135 .unwrap();
136
137 let c1 = SkeletonNodeBuilder::default()
138 .id(1_u16)
139 .name("c1")
140 .original_key("c1")
141 .build(NodeKind::LeafHolder)
142 .unwrap();
143
144 let c2 = SkeletonNodeBuilder::default()
145 .id(2_u16)
146 .name("c2")
147 .original_key("c2")
148 .build(NodeKind::LeafHolder)
149 .unwrap();
150
151 let skel = SkeletonBuilder::default()
152 .nodes(vec![root,c1,c2])
153 .root_id(Some(0_u16)).build().unwrap();
154
155 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
156 let probs = cfg.leaf_probability_per_level();
157 assert_eq!(probs, &[0.0, 1.0]);
158 }
159
160 #[traced_test]
162 fn chain_with_no_structural_leaf() {
163 let n0 = SkeletonNodeBuilder::default()
164 .id(0_u16)
165 .child_ids(vec![1_u16])
166 .name("n0")
167 .original_key("n0")
168 .build(NodeKind::Dispatch)
169 .unwrap();
170
171 let n1 = SkeletonNodeBuilder::default()
172 .id(1_u16)
173 .child_ids(vec![2_u16])
174 .name("n1")
175 .original_key("n1")
176 .build(NodeKind::Dispatch)
177 .unwrap();
178
179 let n2 = SkeletonNodeBuilder::default()
180 .id(2_u16)
181 .child_ids(vec![99_u16])
182 .name("n2")
183 .original_key("n2")
184 .build(NodeKind::Dispatch)
185 .unwrap();
186
187 let skel = SkeletonBuilder::default()
188 .nodes(vec![n0,n1,n2])
189 .root_id(Some(0_u16)).build().unwrap();
190
191 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
192 let probs = cfg.leaf_probability_per_level();
193 assert_eq!(probs.len(), 3);
194 assert!(probs.iter().all(|&p| (p - 0.0).abs() < 1e-6));
195 }
196
197 #[traced_test]
199 fn disconnected_leaf_ignored() {
200 let root = SkeletonNodeBuilder::default()
201 .id(0_u16)
202 .child_ids(vec![1_u16])
203 .name("root")
204 .original_key("root")
205 .build(NodeKind::Dispatch)
206 .unwrap();
207
208 let leaf = SkeletonNodeBuilder::default()
209 .id(1_u16)
210 .name("leaf")
211 .original_key("leaf")
212 .build(NodeKind::LeafHolder)
213 .unwrap();
214
215 let orphan = SkeletonNodeBuilder::default()
216 .id(2_u16)
217 .name("orphan")
218 .original_key("orphan")
219 .build(NodeKind::LeafHolder)
220 .unwrap();
221
222 let skel = SkeletonBuilder::default()
223 .nodes(vec![root,leaf,orphan])
224 .root_id(Some(0_u16))
225 .build()
226 .unwrap();
227
228 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
229 let probs = cfg.leaf_probability_per_level();
230 assert_eq!(probs, &[0.0, 1.0]);
231 }
232
233 #[traced_test]
235 fn fractional_skip_probability() {
236
237 let root = SkeletonNodeBuilder::default()
238 .id(0_u16)
239 .child_ids(vec![1_u16,2_u16,3_u16])
240 .name("root")
241 .original_key("root")
242 .build(NodeKind::Dispatch)
243 .unwrap();
244
245 let leaf1 = SkeletonNodeBuilder::default()
246 .id(1_u16)
247 .name("l1")
248 .original_key("l1")
249 .build(NodeKind::LeafHolder)
250 .unwrap();
251
252 let leaf2 = SkeletonNodeBuilder::default()
253 .id(2_u16)
254 .name("l2")
255 .original_key("l2")
256 .build(NodeKind::LeafHolder)
257 .unwrap();
258
259 let mid = SkeletonNodeBuilder::default()
260 .id(3_u16)
261 .child_ids(vec![4_u16])
262 .name("m")
263 .original_key("m")
264 .build(NodeKind::Dispatch)
265 .unwrap();
266
267 let leaf3 = SkeletonNodeBuilder::default()
268 .id(4_u16)
269 .name("l3")
270 .original_key("l3")
271 .build(NodeKind::LeafHolder)
272 .unwrap();
273
274 let skel = SkeletonBuilder::default()
275 .nodes(vec![root,leaf1,leaf2,mid,leaf3])
276 .root_id(Some(0_u16))
277 .build()
278 .unwrap();
279
280 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
281 let probs = cfg.leaf_probability_per_level();
282 assert_eq!(probs.len(), 3);
283 assert!((probs[0] - 0.0).abs() < 1e-6);
284 let expected = (2.0_f32 / 3.0).powf(1.2);
285 assert!((probs[1] - expected).abs() < 1e-5);
286 assert!((probs[2] - 1.0).abs() < 1e-6);
287 }
288
289 #[traced_test]
295 fn cycle_does_not_loop_and_counts_correctly() {
296 let n0 = SkeletonNodeBuilder::default()
297 .id(0_u16)
298 .child_ids(vec![1_u16])
299 .name("n0")
300 .original_key("n0")
301 .build(NodeKind::Dispatch)
302 .unwrap();
303
304 let n1 = SkeletonNodeBuilder::default()
305 .id(1_u16)
306 .child_ids(vec![0_u16])
307 .name("n1")
308 .original_key("n1")
309 .build(NodeKind::Dispatch)
310 .unwrap();
311
312 let skel = SkeletonBuilder::default()
313 .nodes(vec![n0,n1])
314 .root_id(Some(0_u16)).build().unwrap();
315
316 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
317 let probs = cfg.leaf_probability_per_level();
318 assert_eq!(probs, &[0.0,0.0]);
319 }
320
321 #[traced_test]
323 fn deep_chain_saturates_and_last_is_full() {
324 let nodes: Vec<SkeletonNode> = (0_u16..300).map(|i| {
325 let mut b = SkeletonNodeBuilder::default()
326 .id(i)
327 .name(format!("n{}", i))
328 .original_key(format!("n{}", i));
329 if i < 299 {
330 b = b.child_ids(vec![(i + 1)]);
331 }
332
333 b
334 .build(NodeKind::Dispatch)
335 .unwrap()
336
337 }).collect();
338 let skel = SkeletonBuilder::default()
339 .nodes(nodes)
340 .root_id(Some(0_u16)).build().unwrap();
341
342 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
343 let probs = cfg.leaf_probability_per_level();
344 assert_eq!(probs.len(), 256);
345 assert!((probs[255] - 1.0).abs() < 1e-6);
346 }
347
348 #[traced_test]
350 fn uniform_invalid_children_skipping() {
351 let nodes: Vec<SkeletonNode> = (0_u16..10).map(|i| {
352 let children: Vec<u16> = (0_u16..5_u16).map(|c| 100_u16 + c).collect();
353 SkeletonNodeBuilder::default()
354 .id(i)
355 .child_ids(children)
356 .name(format!("n{}", i))
357 .original_key(format!("n{}", i))
358 .build(NodeKind::Dispatch)
359 .unwrap()
360 }).collect();
361 let skel = SkeletonBuilder::default()
362 .nodes(nodes)
363 .root_id(Some(0_u16)).build().unwrap();
364
365 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
366 let probs = cfg.leaf_probability_per_level();
367 assert_eq!(probs, &[0.0]);
368 }
369
370 #[traced_test]
372 fn leaf_count_field_does_not_affect_skipping() {
373
374 let n0 = SkeletonNodeBuilder::default()
375 .id(0_u16)
376 .child_ids(vec![1_u16])
377 .leaf_count(1000_u16)
378 .name("n0")
379 .original_key("n0")
380 .build(NodeKind::Dispatch)
381 .unwrap();
382
383 let n1 = SkeletonNodeBuilder::default()
384 .id(1_u16)
385 .leaf_count(0_u16)
386 .name("n1")
387 .original_key("n1")
388 .build(NodeKind::LeafHolder)
389 .unwrap();
390
391 let skel = SkeletonBuilder::default()
392 .nodes(vec![n0,n1])
393 .root_id(Some(0_u16))
394 .build()
395 .unwrap();
396
397 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
398 let probs = cfg.leaf_probability_per_level();
399 assert_eq!(probs, &[0.0, 1.0]);
400 }
401
402 #[traced_test]
403 fn all_zero_stats_returns_some() {
404 let n0 = SkeletonNodeBuilder::default()
407 .id(0_u16)
408 .child_ids(vec![1_u16])
409 .name("n0")
410 .original_key("n0")
411 .build(NodeKind::Dispatch)
412 .unwrap();
413
414 let n1 = SkeletonNodeBuilder::default()
415 .id(1_u16)
416 .child_ids(vec![2_u16])
417 .name("n1")
418 .original_key("n1")
419 .build(NodeKind::Dispatch)
420 .unwrap();
421
422 let n2 = SkeletonNodeBuilder::default()
423 .id(2_u16)
424 .child_ids(vec![99_u16]) .name("n2")
426 .original_key("n2")
427 .build(NodeKind::Dispatch)
428 .unwrap();
429
430 let skel = SkeletonBuilder::default()
431 .nodes(vec![n0, n1, n2])
432 .root_id(Some(0_u16))
433 .build()
434 .unwrap();
435
436 let cfg = skel.try_infer_level_skipping_configuration().unwrap();
438 let probs = cfg.leaf_probability_per_level();
439 assert!(probs.iter().all(|&p| (p - 0.0).abs() < 1e-6),
440 "Expected all skip probabilities to be 0.0 when there are no valid structural leaves");
441 }
442}