1//! Choice over more than 255 labels by walking a label tree (contract 2//! *SQL surface*, "More than 255 options", label tree mode). 3//! 4//! Each internal node is one Choice over its children. A node's score is 5//! the product of the edge probabilities from the root, kept as 6//! `Σ ln max(p, 1e-9)`. That is a real probability, and it can only fall 7//! with depth, so a scored node bounds every leaf below it. The 8//! cookbook's geometric mean `exp(mean(ln p))` is not used: it rises with 9//! depth and prefers deep paths of confident edges over a shallow likely 10//! one. 11//! 12//! The search is best-first, K nodes per round (K = 3 by default): each 13//! round asks the Choices of the K best unexpanded nodes, which the 14//! caller sends in one request. Nothing is dropped for being outside the 15//! K; a node is dropped only when its bound shows it cannot change the 16//! result. So the answer is exact: the best leaf, and the runner-up the 17//! separation `exp(top − second)` is measured against. 18//! 19//! With `tau`, the search returns the deepest node whose path probability 20//! is at least τ (ties at one depth to the likelier), and never expands a 21//! node below τ. Since sibling probabilities share their parent's, at 22//! most `1/τ` nodes per depth qualify. 23//! 24//! A node with one child is not a question (`jev_protocol::Choice`): its 25//! edge is ×1 and it is passed through without a request. 26//! 27//! Two things ride in a round beside its K nodes, and neither changes 28//! the result: 29//! 30//! - **Lookahead** ([`Lookahead`]): the Choices of each asked node's 31//! likeliest children, within a token budget, so one round trip can 32//! score two levels. An answer that arrives before its node is scored 33//! is held, and applied without a request once it is. 34//! - **The order twin** ([`TwinBelow`]): a node whose top two children 35//! come back close is not expanded; it is queued again and asked in 36//! reverse order, and the two answers' log-probabilities are averaged. 37//! Closeness is only known from the answer, so the twin goes in a 38//! later round, in its turn by score like any node: it costs a round 39//! trip only when nothing else is left to ask, and nothing when its 40//! node is pruned first. 41//! 42//! Sans-I/O: [`Search::step`] says which Choices to ask, [`Tree::choice`] 43//! builds each one (options describing their subtree, instructions 44//! naming the path, in [`Ask::order`]), and [`Search::answer`] takes its 45//! probabilities in the order sent. Sending them is the caller's. 46 47use std::collections::{HashMap, HashSet}; 48use std::num::NonZeroUsize; 49 50use jev_protocol::{Choice, Json, MAX_CHOICE_OPTIONS, Noul}; 51 52/// The floor under an edge probability, so a 0 costs `ln 1e-9` rather 53/// than sending a path to −∞ and making every such path tie. 54pub const PROBABILITY_FLOOR: f64 = 1e-9; 55 56/// The vendor's beam width (hierarchical-classification cookbook); not 57/// measured by us. 58pub const DEFAULT_K: NonZeroUsize = NonZeroUsize::new(3).unwrap(); 59 60/// A node of a [`Tree`], valid only for the tree that issued it. 61#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] 62pub struct NodeId(u32); 63 64/// A label hierarchy as the caller writes it. 65#[derive(Clone, Debug, PartialEq)] 66pub struct Branch { 67 pub label: String, 68 pub children: Vec<Branch>, 69} 70 71impl Branch { 72 pub fn leaf(label: impl Into<String>) -> Self { 73 Branch { label: label.into(), children: Vec::new() } 74 } 75 76 pub fn node(label: impl Into<String>, children: impl IntoIterator<Item = Branch>) -> Self { 77 Branch { label: label.into(), children: children.into_iter().collect() } 78 } 79} 80 81impl Branch { 82 /// A root over `paths`, each a sequence of labels from the top down. 83 /// Siblings keep the order in which their label first appears, since 84 /// a node's options are sent in that order. Every path ends at a leaf, 85 /// so a path that is a prefix of another is refused rather than 86 /// silently dropped from the answers a search can return. 87 pub fn from_paths<P, L>(paths: P) -> Result<Branch, PathError> 88 where 89 P: IntoIterator<Item = L>, 90 L: IntoIterator, 91 L::Item: Into<String>, 92 { 93 let mut root = Branch::leaf(""); 94 let mut ends: HashSet<Vec<String>> = HashSet::new(); 95 for path in paths { 96 let path: Vec<String> = path.into_iter().map(Into::into).collect(); 97 if path.is_empty() { 98 return Err(PathError::Empty); 99 } 100 let mut at = &mut root; 101 for (i, label) in path.iter().enumerate() { 102 if ends.contains(&path[..i]) && i > 0 { 103 return Err(PathError::LeafAndGroup { path: path[..i].to_vec() }); 104 } 105 let c = match at.children.iter().position(|c| c.label == *label) { 106 Some(c) => c, 107 None => { 108 at.children.push(Branch::leaf(label.clone())); 109 at.children.len() - 1 110 } 111 }; 112 at = &mut at.children[c]; 113 } 114 if !at.children.is_empty() { 115 return Err(PathError::LeafAndGroup { path }); 116 } 117 if !ends.insert(path.clone()) { 118 return Err(PathError::Repeated { path }); 119 } 120 } 121 Ok(root) 122 } 123} 124 125/// Paths that do not describe a tree of leaves. 126#[derive(Clone, Debug, PartialEq, Eq)] 127pub enum PathError { 128 Empty, 129 Repeated { path: Vec<String> }, 130 /// A path that also groups other paths below it. 131 LeafAndGroup { path: Vec<String> }, 132} 133 134impl std::fmt::Display for PathError { 135 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { 136 match self { 137 PathError::Empty => write!(f, "a label path is empty"), 138 PathError::Repeated { path } => write!(f, "the label path {path:?} is given twice"), 139 PathError::LeafAndGroup { path } => { 140 write!(f, "the label path {path:?} is both a label and the parent of other labels") 141 } 142 } 143 } 144} 145 146impl std::error::Error for PathError {} 147 148#[derive(Clone, Debug, PartialEq, Eq)] 149pub enum TreeError { 150 /// The root has no children, so there is nothing to choose. 151 NoChoice, 152 /// More children than one Choice can carry. 153 Fanout { path: Vec<String>, children: usize }, 154 EmptyLabel { path: Vec<String> }, 155 DuplicateLabel { path: Vec<String>, label: String }, 156} 157 158impl std::fmt::Display for TreeError { 159 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { 160 match self { 161 TreeError::NoChoice => write!(f, "the label tree has no labels under its root"), 162 TreeError::Fanout { path, children } => write!( 163 f, 164 "the label tree node {path:?} has {children} children; a Choice takes at most {MAX_CHOICE_OPTIONS}" 165 ), 166 TreeError::EmptyLabel { path } => write!(f, "a label under {path:?} is empty"), 167 TreeError::DuplicateLabel { path, label } => { 168 write!(f, "the label {label:?} appears twice under {path:?}") 169 } 170 } 171 } 172} 173 174impl std::error::Error for TreeError {} 175 176#[derive(Clone, Debug)] 177struct Node { 178 label: String, 179 parent: Option<NodeId>, 180 children: Vec<NodeId>, 181 depth: u32, 182 /// Edges from this node down to its deepest leaf. 183 height: u32, 184 /// Leaves under this node; a leaf counts itself. 185 leaves: u32, 186} 187 188/// A validated label hierarchy: every node has at most 255 children with 189/// unique, non-empty labels, so every node's Choice can be built. 190#[derive(Clone, Debug)] 191pub struct Tree { 192 nodes: Vec<Node>, 193} 194 195impl Tree { 196 /// The root's own label is never sent; its children are the first 197 /// Choice. 198 pub fn new(root: Branch) -> Result<Self, TreeError> { 199 if root.children.is_empty() { 200 return Err(TreeError::NoChoice); 201 } 202 let mut tree = Tree { nodes: Vec::new() }; 203 tree.add(root, None, 0)?; 204 Ok(tree) 205 } 206 207 fn add(&mut self, b: Branch, parent: Option<NodeId>, depth: u32) -> Result<NodeId, TreeError> { 208 let id = NodeId(self.nodes.len() as u32); 209 self.nodes.push(Node { label: b.label, parent, children: Vec::new(), depth, height: 0, leaves: 1 }); 210 if b.children.len() > MAX_CHOICE_OPTIONS { 211 return Err(TreeError::Fanout { path: self.path_owned(id), children: b.children.len() }); 212 } 213 let mut seen = HashSet::new(); 214 for c in &b.children { 215 if c.label.is_empty() { 216 return Err(TreeError::EmptyLabel { path: self.path_owned(id) }); 217 } 218 if !seen.insert(c.label.clone()) { 219 return Err(TreeError::DuplicateLabel { path: self.path_owned(id), label: c.label.clone() }); 220 } 221 } 222 let (mut height, mut leaves) = (0, 0); 223 for c in b.children { 224 let child = self.add(c, Some(id), depth + 1)?; 225 self.nodes[id.0 as usize].children.push(child); 226 height = height.max(self.nodes[child.0 as usize].height + 1); 227 leaves += self.nodes[child.0 as usize].leaves; 228 } 229 self.nodes[id.0 as usize].height = height; 230 self.nodes[id.0 as usize].leaves = leaves.max(1); 231 Ok(id) 232 } 233 234 fn node(&self, n: NodeId) -> &Node { 235 &self.nodes[n.0 as usize] 236 } 237 238 pub fn root(&self) -> NodeId { 239 NodeId(0) 240 } 241 242 pub fn label(&self, n: NodeId) -> &str { 243 &self.node(n).label 244 } 245 246 /// In the caller's order, which is the order the Choice sends. 247 pub fn children(&self, n: NodeId) -> &[NodeId] { 248 &self.node(n).children 249 } 250 251 pub fn is_leaf(&self, n: NodeId) -> bool { 252 self.node(n).children.is_empty() 253 } 254 255 /// Edges from the root; the root is 0. 256 pub fn depth(&self, n: NodeId) -> u32 { 257 self.node(n).depth 258 } 259 260 /// The labels from below the root down to `n`. 261 pub fn path(&self, n: NodeId) -> Vec<&str> { 262 let mut path = Vec::new(); 263 let mut at = n; 264 while let Some(parent) = self.node(at).parent { 265 path.push(self.label(at)); 266 at = parent; 267 } 268 path.reverse(); 269 path 270 } 271 272 fn path_owned(&self, n: NodeId) -> Vec<String> { 273 self.path(n).into_iter().map(str::to_owned).collect() 274 } 275 276 /// The leaves under `n`, depth first in the caller's order. 277 fn leaves(&self, n: NodeId) -> Vec<NodeId> { 278 let mut out = Vec::new(); 279 let mut stack = vec![n]; 280 while let Some(at) = stack.pop() { 281 if self.is_leaf(at) { 282 out.push(at); 283 } 284 stack.extend(self.children(at).iter().rev()); 285 } 286 out 287 } 288 289 /// `n`'s Choice over its children, in `order` (the caller's, or 290 /// reversed for an order-debiasing twin), each label the option key 291 /// (contract *SQL surface*). The instructions are 292 /// `question` and, below the root, the path to `n`; an internal 293 /// child's description lists its direct children and a sample of the 294 /// leaves under it, so the model sees what lives under a branch 295 /// (vendor guidance; its hierarchical cookbook does neither). A leaf 296 /// child has no description. `None` for a node with fewer than two 297 /// children, which [`Search`] never asks. 298 /// 299 /// Every piece is written in a fixed form, so equal nodes give equal 300 /// bytes and one cache key. 301 pub fn choice(&self, n: NodeId, order: Order, question: &str, describe: Describe) -> Option<Choice> { 302 let children = self.children(n); 303 if children.len() < 2 { 304 return None; 305 } 306 let path = self.path(n); 307 let instructions = if path.is_empty() { 308 question.to_owned() 309 } else { 310 format!("{question}\n\nWithin: {}", path.join(" > ")) 311 }; 312 let option = |&c: &NodeId| (self.label(c).to_owned(), self.describe(c, describe).map(|d| Json::text(&d))); 313 let options: Vec<_> = match order { 314 Order::Caller => children.iter().map(option).collect(), 315 Order::Reversed => children.iter().rev().map(option).collect(), 316 }; 317 Some(Choice::new(Json::text(&instructions), options).expect("a Tree node's children make a valid Choice")) 318 } 319 320 /// The "does any label fit" Noul ([`crate::gate`]), with the tree shown 321 /// as the root is described: its top-level labels and a sample of 322 /// leaves, so its size does not grow with the tree. `None` for a tree 323 /// with one leaf, which the search answers without a request. 324 pub fn gate(&self, question: &str, describe: Describe) -> Option<Noul> { 325 let root = self.root(); 326 if self.leaves(root).len() < 2 { 327 return None; 328 } 329 let labels = self.describe(root, describe).expect("a root with children is described"); 330 Some(crate::gate::gate(question, &labels)) 331 } 332 333 fn describe(&self, n: NodeId, describe: Describe) -> Option<String> { 334 let children = self.children(n); 335 if children.is_empty() { 336 return None; 337 } 338 let listed = children.len().min(describe.children); 339 let mut text = format!( 340 "Contains: {}", 341 children[..listed].iter().map(|&c| self.label(c)).collect::<Vec<_>>().join(", ") 342 ); 343 if listed < children.len() { 344 text.push_str(&format!(", and {} more", children.len() - listed)); 345 } 346 // Leaves that are not already listed as direct children. 347 let sample: Vec<&str> = self 348 .leaves(n) 349 .into_iter() 350 .filter(|l| self.node(*l).parent != Some(n)) 351 .take(describe.leaves) 352 .map(|l| self.label(l)) 353 .collect(); 354 if !sample.is_empty() { 355 text.push_str(&format!(". For example: {}", sample.join(", "))); 356 } 357 Some(text) 358 } 359} 360 361/// How much of a subtree an option's description shows. The defaults 362/// are ours and unmeasured; the label-tree release gate (*Testing*) 363/// measures them with K. 364#[derive(Clone, Copy, Debug, PartialEq, Eq)] 365pub struct Describe { 366 /// Direct children listed, in the caller's order. 367 pub children: usize, 368 /// Leaves below those children, depth first. 369 pub leaves: usize, 370} 371 372impl Default for Describe { 373 fn default() -> Self { 374 Describe { children: 20, leaves: 5 } 375 } 376} 377 378/// The order a Choice's options are sent in. 379#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] 380pub enum Order { 381 Caller, 382 /// The order-debiasing twin's. 383 Reversed, 384} 385 386/// A probability threshold for early stop, in (0, 1]. 387#[derive(Clone, Copy, Debug, PartialEq)] 388pub struct Tau(f64); 389 390impl Tau { 391 pub fn new(tau: f64) -> Option<Self> { 392 (tau > 0.0 && tau <= 1.0).then_some(Tau(tau)) 393 } 394 395 pub fn get(self) -> f64 { 396 self.0 397 } 398} 399 400/// Speculative Choices carried in a round: for each node asked, the 401/// Choices of its `m` likeliest children, while the round's estimated 402/// tokens stay within `budget`. So one round trip scores two levels. 403/// 404/// "Likeliest" is judged before the parent is answered, since both ride 405/// in one request, so it is a prior: the children with the most leaves 406/// under them (a uniform prior over leaves), ties in the caller's order. 407/// `m` and the prior are unmeasured; the label-tree release gate 408/// (*Testing*) measures them. 409#[derive(Clone, Copy, Debug, PartialEq, Eq)] 410pub struct Lookahead { 411 pub m: usize, 412 /// Estimated input tokens for the whole round, the asked nodes' 413 /// Choices included. The asked nodes are sent even when they alone 414 /// exceed it. 415 pub budget: usize, 416} 417 418impl Default for Lookahead { 419 /// Below jev-1.13.0's 64k tokens per request, leaving room for the 420 /// request's own overhead. 421 fn default() -> Self { 422 Lookahead { m: 2, budget: 60_000 } 423 } 424} 425 426/// A node whose top two children are within this separation 427/// (`p_top / p_second`, floored) is a close call, and its Choice is asked 428/// again in reverse order; the two answers' log-probabilities are 429/// averaged (permutation self-consistency, arXiv 2310.07712), since the 430/// measured order bias favours whatever is listed last. Greater than 1. 431#[derive(Clone, Copy, Debug, PartialEq)] 432pub struct TwinBelow(f64); 433 434impl TwinBelow { 435 pub fn new(separation: f64) -> Option<Self> { 436 (separation > 1.0 && separation.is_finite()).then_some(TwinBelow(separation)) 437 } 438 439 pub fn get(self) -> f64 { 440 self.0 441 } 442} 443 444#[derive(Clone, Copy, Debug, PartialEq)] 445pub struct Params { 446 /// Scored nodes asked per round, all in one request. A twin takes a 447 /// slot, since it is its node's own question; lookahead Choices ride 448 /// outside it. 449 pub k: NonZeroUsize, 450 /// Unset: always return a leaf. 451 pub tau: Option<Tau>, 452 /// Unset: no speculative Choices. 453 pub lookahead: Option<Lookahead>, 454 /// Unset: no twins. 455 pub twin: Option<TwinBelow>, 456} 457 458impl Default for Params { 459 /// The twin's 1.5× is ours and unmeasured, like [`Lookahead`]'s 460 /// defaults. 461 fn default() -> Self { 462 Params { k: DEFAULT_K, tau: None, lookahead: Some(Lookahead::default()), twin: TwinBelow::new(1.5) } 463 } 464} 465 466/// One Choice a round asks. Build it with [`Tree::choice`] in 467/// [`Ask::order`], and answer it with [`Search::answer`]. 468#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] 469pub enum Ask { 470 /// A scored node, best first. 471 Node(NodeId), 472 /// A close call's second asking, in reverse order. 473 Twin(NodeId), 474 /// A child of a node asked this round, before its score is known. 475 Lookahead(NodeId), 476} 477 478impl Ask { 479 pub fn node(self) -> NodeId { 480 match self { 481 Ask::Node(n) | Ask::Twin(n) | Ask::Lookahead(n) => n, 482 } 483 } 484 485 pub fn order(self) -> Order { 486 match self { 487 Ask::Twin(_) => Order::Reversed, 488 Ask::Node(_) | Ask::Lookahead(_) => Order::Caller, 489 } 490 } 491} 492 493/// What the search needs next. 494#[derive(Clone, Debug, PartialEq)] 495pub enum Step { 496 /// Ask these Choices, in one request, and [`Search::answer`] each. 497 Ask(Vec<Ask>), 498 Done(Outcome), 499} 500 501#[derive(Clone, Debug, PartialEq)] 502pub enum Outcome { 503 /// The most probable leaf. 504 Leaf { 505 leaf: NodeId, 506 probability: f64, 507 /// `exp(score_top − score_second)`: near 1 is ambiguous. `None` 508 /// when the tree admits one leaf only. 509 separation: Option<f64>, 510 }, 511 /// Under `tau`: the deepest node at or above τ. 512 Stop { node: NodeId, depth: u32, leaf: bool, probability: f64 }, 513} 514 515#[derive(Clone, Debug, PartialEq)] 516pub enum AnswerError { 517 /// Not asked in the current round, or already answered. 518 NotAsked(Ask), 519 /// One probability per option, in the order sent. 520 Count { node: NodeId, expected: usize, got: usize }, 521 NotAProbability { node: NodeId, value: f64 }, 522} 523 524impl std::fmt::Display for AnswerError { 525 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { 526 match self { 527 AnswerError::NotAsked(a) => write!(f, "label tree Choice {a:?} was not asked"), 528 AnswerError::Count { node, expected, got } => { 529 write!(f, "label tree node {node:?} has {expected} options, answered with {got}") 530 } 531 AnswerError::NotAProbability { node, value } => { 532 write!(f, "label tree node {node:?} was answered with {value}, not a probability") 533 } 534 } 535 } 536} 537 538impl std::error::Error for AnswerError {} 539 540#[derive(Clone, Copy, Debug)] 541struct Entry { 542 score: f64, 543 node: NodeId, 544} 545 546/// One label-tree search: best-first, exact, K nodes per round. 547/// 548/// Lookahead and twins change what a round carries, never the result: 549/// an answer that arrives early is held until its node is scored, and 550/// the node is then expanded without a request. The pruning is the same 551/// bound either way. 552#[derive(Debug)] 553pub struct Search<'t> { 554 tree: &'t Tree, 555 params: Params, 556 /// Scored and not yet expanded or finalized, best first. A queued 557 /// node with an entry in `answers` is a close call awaiting its twin. 558 queue: Vec<Entry>, 559 /// Asked this round; `score` is `None` for a lookahead, whose node is 560 /// not scored yet. 561 pending: Vec<(Ask, Option<f64>)>, 562 /// Answers not yet applied, in child order: a lookahead's, until its 563 /// node is scored, and a close call's first, until its twin. 564 answers: HashMap<NodeId, Vec<f64>>, 565 /// Leaves proven best, in order (no `tau`). 566 finalized: Vec<Entry>, 567 /// The deepest node at or above τ so far (`tau`). 568 stop: Option<Entry>, 569} 570 571impl<'t> Search<'t> { 572 pub fn new(tree: &'t Tree, params: Params) -> Self { 573 let mut s = Search { 574 tree, 575 params, 576 queue: Vec::new(), 577 pending: Vec::new(), 578 answers: HashMap::new(), 579 finalized: Vec::new(), 580 stop: None, 581 }; 582 s.push(Entry { score: 0.0, node: tree.root() }); 583 s 584 } 585 586 fn ln_tau(&self) -> Option<f64> { 587 self.params.tau.map(|t| t.get().ln()) 588 } 589 590 /// Scores a node, passing through one-child nodes at ×1, and expands 591 /// it at once if its answer is already held. 592 fn push(&mut self, mut e: Entry) { 593 loop { 594 if let Some(ln_tau) = self.ln_tau() { 595 if e.score < ln_tau { 596 return; 597 } 598 if self.better_stop(e) { 599 self.stop = Some(e); 600 } 601 } 602 match self.tree.children(e.node) { 603 [only] => e.node = *only, 604 _ => break, 605 } 606 } 607 if self.params.tau.is_some() && self.tree.is_leaf(e.node) { 608 return; 609 } 610 if let Some(first) = self.answers.get(&e.node) 611 && !self.is_close(first) 612 { 613 let first = self.answers.remove(&e.node).expect("just read"); 614 let ln: Vec<f64> = first.iter().map(|p| ln_floored(*p)).collect(); 615 return self.expand(e, &ln); 616 } 617 self.enqueue(e); 618 } 619 620 fn enqueue(&mut self, e: Entry) { 621 let at = self.queue.partition_point(|q| rank(q, &e).is_lt()); 622 self.queue.insert(at, e); 623 } 624 625 fn expand(&mut self, parent: Entry, ln: &[f64]) { 626 let children = self.tree.children(parent.node); 627 for (&child, &l) in children.iter().zip(ln) { 628 self.push(Entry { score: parent.score + l, node: child }); 629 } 630 } 631 632 /// Whether this first answer, in child order, needs a twin. 633 fn is_close(&self, probabilities: &[f64]) -> bool { 634 let Some(twin) = self.params.twin else { return false }; 635 let (mut top, mut second) = (0.0f64, 0.0f64); 636 for &p in probabilities { 637 let p = p.max(PROBABILITY_FLOOR); 638 if p > top { 639 (top, second) = (p, top); 640 } else if p > second { 641 second = p; 642 } 643 } 644 top / second <= twin.get() 645 } 646 647 fn better_stop(&self, e: Entry) -> bool { 648 self.stop.is_none_or(|s| { 649 let (d, sd) = (self.tree.depth(e.node), self.tree.depth(s.node)); 650 d > sd || (d == sd && e.score > s.score) 651 }) 652 } 653 654 /// Whether anything under `e` could replace the current stop. 655 fn could_beat_stop(&self, e: &Entry) -> bool { 656 self.stop.is_none_or(|s| { 657 let reach = self.tree.depth(e.node) + self.tree.node(e.node).height; 658 let sd = self.tree.depth(s.node); 659 reach > sd || (reach == sd && e.score > s.score) 660 }) 661 } 662 663 fn round(&self) -> Step { 664 Step::Ask(self.pending.iter().map(|(a, _)| *a).collect()) 665 } 666 667 /// The next round to ask, or the result. Asking again before every 668 /// asked Choice is answered returns the same round. 669 /// 670 /// `tokens` estimates one Choice's input tokens, for the lookahead 671 /// budget; it is not called without [`Params::lookahead`]. 672 pub fn step(&mut self, tokens: impl FnMut(Ask) -> usize) -> Step { 673 if !self.pending.is_empty() { 674 return self.round(); 675 } 676 if self.params.tau.is_some() { 677 return self.step_tau(tokens); 678 } 679 // A leaf at the head of the queue beats every leaf still unscored, 680 // because scores only fall with depth. 681 while let Some(head) = self.queue.first() { 682 if !self.tree.is_leaf(head.node) { 683 break; 684 } 685 self.finalized.push(self.queue.remove(0)); 686 if self.finalized.len() == 2 { 687 return Step::Done(self.leaf_outcome()); 688 } 689 } 690 if self.queue.is_empty() { 691 return Step::Done(self.leaf_outcome()); 692 } 693 // Only the best two leaves matter: a node no likelier than the 694 // second-best known leaf cannot change them. 695 let needed = 2 - self.finalized.len(); 696 let bound = self.queue.iter().filter(|e| self.tree.is_leaf(e.node)).nth(needed - 1).map(|e| e.score); 697 if let Some(bound) = bound { 698 let tree = self.tree; 699 self.queue.retain(|e| tree.is_leaf(e.node) || e.score > bound); 700 } 701 self.take_round(|tree, e| !tree.is_leaf(e.node), tokens) 702 } 703 704 fn step_tau(&mut self, tokens: impl FnMut(Ask) -> usize) -> Step { 705 let queue = std::mem::take(&mut self.queue); 706 self.queue = queue.into_iter().filter(|e| self.could_beat_stop(e)).collect(); 707 if self.queue.is_empty() { 708 let s = self.stop.expect("the root is always at or above τ"); 709 return Step::Done(Outcome::Stop { 710 node: s.node, 711 depth: self.tree.depth(s.node), 712 leaf: self.tree.is_leaf(s.node), 713 probability: s.score.exp(), 714 }); 715 } 716 self.take_round(|_, _| true, tokens) 717 } 718 719 fn take_round(&mut self, askable: impl Fn(&Tree, &Entry) -> bool, mut tokens: impl FnMut(Ask) -> usize) -> Step { 720 let k = self.params.k.get(); 721 let mut i = 0; 722 while i < self.queue.len() && self.pending.len() < k { 723 if askable(self.tree, &self.queue[i]) { 724 let e = self.queue.remove(i); 725 let ask = if self.answers.contains_key(&e.node) { Ask::Twin(e.node) } else { Ask::Node(e.node) }; 726 self.pending.push((ask, Some(e.score))); 727 } else { 728 i += 1; 729 } 730 } 731 if let Some(look) = self.params.lookahead { 732 self.add_lookahead(look, &mut tokens); 733 } 734 self.round() 735 } 736 737 /// Fills the round with the asked nodes' likeliest children's Choices, 738 /// best-ranked asked node first, within the budget. 739 fn add_lookahead(&mut self, look: Lookahead, tokens: &mut impl FnMut(Ask) -> usize) { 740 let mut used: usize = self.pending.iter().map(|(a, _)| tokens(*a)).sum(); 741 let parents: Vec<NodeId> = self.pending.iter().map(|(a, _)| a.node()).collect(); 742 for parent in parents { 743 let mut children = self.tree.children(parent).to_vec(); 744 // Stable, so ties keep the caller's order. 745 children.sort_by_key(|&c| std::cmp::Reverse(self.tree.node(c).leaves)); 746 for c in children.into_iter().take(look.m) { 747 let Some(n) = self.askable_below(c) else { continue }; 748 if self.answers.contains_key(&n) || self.pending.iter().any(|(a, _)| a.node() == n) { 749 continue; 750 } 751 let ask = Ask::Lookahead(n); 752 let cost = tokens(ask); 753 if used + cost > look.budget { 754 continue; 755 } 756 used += cost; 757 self.pending.push((ask, None)); 758 } 759 } 760 } 761 762 /// The node a Choice would be asked for once `n` is scored: `n` 763 /// passed through its one-child chain, or `None` at a leaf. 764 fn askable_below(&self, mut n: NodeId) -> Option<NodeId> { 765 loop { 766 match self.tree.children(n) { 767 [] => return None, 768 [only] => n = *only, 769 _ => return Some(n), 770 } 771 } 772 } 773 774 fn leaf_outcome(&self) -> Outcome { 775 let top = self.finalized[0]; 776 Outcome::Leaf { 777 leaf: top.node, 778 probability: top.score.exp(), 779 separation: self.finalized.get(1).map(|second| (top.score - second.score).exp()), 780 } 781 } 782 783 /// A Choice's answer: one probability per option, in the order it was 784 /// sent ([`Ask::order`]). 785 pub fn answer(&mut self, ask: Ask, probabilities: &[f64]) -> Result<(), AnswerError> { 786 let at = self.pending.iter().position(|(a, _)| *a == ask).ok_or(AnswerError::NotAsked(ask))?; 787 let node = ask.node(); 788 let children = self.tree.children(node); 789 if probabilities.len() != children.len() { 790 return Err(AnswerError::Count { node, expected: children.len(), got: probabilities.len() }); 791 } 792 if let Some(&value) = probabilities.iter().find(|p| !(0.0..=1.0).contains(*p)) { 793 return Err(AnswerError::NotAProbability { node, value }); 794 } 795 let mut in_child_order = probabilities.to_vec(); 796 if ask.order() == Order::Reversed { 797 in_child_order.reverse(); 798 } 799 let (_, score) = self.pending.remove(at); 800 match ask { 801 Ask::Node(_) => { 802 let e = Entry { score: score.expect("an asked node is scored"), node }; 803 if self.is_close(&in_child_order) { 804 self.answers.insert(node, in_child_order); 805 self.enqueue(e); 806 } else { 807 let ln: Vec<f64> = in_child_order.iter().map(|p| ln_floored(*p)).collect(); 808 self.expand(e, &ln); 809 } 810 } 811 Ask::Twin(_) => { 812 let e = Entry { score: score.expect("a twin's node is scored"), node }; 813 let first = self.answers.remove(&node).expect("a twin follows a first answer"); 814 let ln: Vec<f64> = 815 first.iter().zip(&in_child_order).map(|(a, b)| (ln_floored(*a) + ln_floored(*b)) / 2.0).collect(); 816 self.expand(e, &ln); 817 } 818 Ask::Lookahead(_) => { 819 self.answers.insert(node, in_child_order); 820 // Its parent may have been answered first in this round. 821 if let Some(i) = self.queue.iter().position(|e| e.node == node) { 822 let e = self.queue.remove(i); 823 self.push(e); 824 } 825 } 826 } 827 Ok(()) 828 } 829} 830 831fn ln_floored(p: f64) -> f64 { 832 p.max(PROBABILITY_FLOOR).ln() 833} 834 835/// Best first; equal scores in tree order, so a search is deterministic. 836fn rank(a: &Entry, b: &Entry) -> std::cmp::Ordering { 837 b.score.total_cmp(&a.score).then(a.node.cmp(&b.node)) 838} 839 840#[cfg(test)] 841mod tests { 842 use super::*; 843 844 /// Runs a search against a fixed answer per internal node, by path 845 /// and in child order (reversed for a twin), and returns the rounds. 846 fn run_rounds( 847 tree: &Tree, 848 params: Params, 849 answer: impl Fn(&[&str], Order) -> Vec<f64>, 850 ) -> (Outcome, Vec<Vec<Ask>>) { 851 let mut s = Search::new(tree, params); 852 let mut rounds = Vec::new(); 853 loop { 854 match s.step(|_| 100) { 855 Step::Done(o) => return (o, rounds), 856 Step::Ask(round) => { 857 assert!(!round.is_empty()); 858 assert!(round.iter().filter(|a| !matches!(a, Ask::Lookahead(_))).count() <= params.k.get()); 859 // Lookahead first, so a held answer is applied when its 860 // parent's lands, as well as the other way round. 861 for &a in round.iter().rev() { 862 s.answer(a, &answer(&tree.path(a.node()), a.order())).unwrap(); 863 } 864 rounds.push(round); 865 } 866 } 867 } 868 } 869 870 /// An answer by path that is the same in either order. 871 fn fixed(answer: impl Fn(&[&str]) -> Vec<f64>) -> impl Fn(&[&str], Order) -> Vec<f64> { 872 move |path, order| { 873 let mut p = answer(path); 874 if order == Order::Reversed { 875 p.reverse(); 876 } 877 p 878 } 879 } 880 881 /// The search alone: no lookahead, no twins. 882 fn plain() -> Params { 883 Params { k: DEFAULT_K, tau: None, lookahead: None, twin: None } 884 } 885 886 #[test] 887 fn the_gate_shows_the_tree_as_the_root_is_described() { 888 let tree = Tree::new(Branch::from_paths([vec!["A", "a1"], vec!["A", "a2"], vec!["B"]]).unwrap()).unwrap(); 889 let want = crate::gate::gate("q", "Contains: A, B. For example: a1, a2"); 890 assert_eq!(tree.gate("q", Describe::default()), Some(want)); 891 let one = Tree::new(Branch::from_paths([vec!["A"]]).unwrap()).unwrap(); 892 assert_eq!(one.gate("q", Describe::default()), None); 893 } 894 895 #[test] 896 fn from_paths_groups_by_prefix_in_first_order() { 897 let root = Branch::from_paths([vec!["B", "y"], vec!["A"], vec!["B", "x"], vec!["C", "z", "1"]]).unwrap(); 898 let want = Branch::node( 899 "", 900 [ 901 Branch::node("B", [Branch::leaf("y"), Branch::leaf("x")]), 902 Branch::leaf("A"), 903 Branch::node("C", [Branch::node("z", [Branch::leaf("1")])]), 904 ], 905 ); 906 assert_eq!(root, want); 907 assert!(Tree::new(root).is_ok()); 908 } 909 910 #[test] 911 fn from_paths_refuses_what_is_not_a_tree_of_leaves() { 912 let none: [Vec<&str>; 1] = [vec![]]; 913 assert_eq!(Branch::from_paths(none), Err(PathError::Empty)); 914 assert_eq!( 915 Branch::from_paths([vec!["A", "x"], vec!["A", "x"]]), 916 Err(PathError::Repeated { path: vec!["A".into(), "x".into()] }) 917 ); 918 // Either order: the group first, or the leaf first. 919 assert_eq!( 920 Branch::from_paths([vec!["A", "x"], vec!["A"]]), 921 Err(PathError::LeafAndGroup { path: vec!["A".into()] }) 922 ); 923 assert_eq!( 924 Branch::from_paths([vec!["A"], vec!["A", "x"]]), 925 Err(PathError::LeafAndGroup { path: vec!["A".into()] }) 926 ); 927 } 928 929 /// With [`plain`] params, the Choices asked. 930 fn run(tree: &Tree, params: Params, answer: impl Fn(&[&str]) -> Vec<f64>) -> (Outcome, usize) { 931 let (o, rounds) = run_rounds(tree, params, fixed(answer)); 932 (o, rounds.iter().map(Vec::len).sum()) 933 } 934 935 fn chain(label: &str, depth: usize, leaf: &str) -> Branch { 936 (0..depth).rev().fold(Branch::leaf(leaf), |below, i| { 937 Branch::node(format!("{label}{i}"), [below, Branch::leaf(format!("{label}{i}-other"))]) 938 }) 939 } 940 941 fn leaf_of(o: &Outcome) -> NodeId { 942 match o { 943 Outcome::Leaf { leaf, .. } => *leaf, 944 other => panic!("{other:?}"), 945 } 946 } 947 948 fn sent(c: &Choice) -> String { 949 jev_protocol::question_bytes(&jev_protocol::Question::Choice(c.clone())) 950 .map(|b| String::from_utf8(b).unwrap()) 951 .unwrap() 952 } 953 954 /// Options describe their subtree; instructions name the path; a leaf 955 /// has no description; a one-child node has no Choice. 956 #[test] 957 fn node_choices() { 958 let tree = Tree::new(Branch::node( 959 "root", 960 [ 961 Branch::node("Animals", [ 962 Branch::node("Birds", [Branch::leaf("Owl"), Branch::leaf("Wren")]), 963 Branch::leaf("Fish"), 964 ]), 965 Branch::leaf("Plants"), 966 Branch::node("Rocks", [Branch::leaf("Granite")]), 967 ], 968 )) 969 .unwrap(); 970 let root = tree.root(); 971 let c = tree.choice(root, Order::Caller, "What is it?", Describe::default()).unwrap(); 972 assert_eq!( 973 sent(&c), 974 r#"{"type":"choice","instructions":"What is it?","criteria":{"Animals":"Contains: Birds, Fish. For example: Owl, Wren","Plants":null,"Rocks":"Contains: Granite"}}"# 975 ); 976 let twin = tree.choice(root, Order::Reversed, "What is it?", Describe::default()).unwrap(); 977 assert_eq!( 978 sent(&twin), 979 r#"{"type":"choice","instructions":"What is it?","criteria":{"Rocks":"Contains: Granite","Plants":null,"Animals":"Contains: Birds, Fish. For example: Owl, Wren"}}"# 980 ); 981 let animals = tree.children(root)[0]; 982 let birds = tree.children(animals)[0]; 983 let c = tree.choice(birds, Order::Caller, "What is it?", Describe::default()).unwrap(); 984 assert_eq!( 985 sent(&c), 986 r#"{"type":"choice","instructions":"What is it?\n\nWithin: Animals > Birds","criteria":{"Owl":null,"Wren":null}}"# 987 ); 988 let rocks = tree.children(root)[2]; 989 assert!(tree.choice(rocks, Order::Caller, "q", Describe::default()).is_none()); 990 let short = tree.choice(root, Order::Caller, "q", Describe { children: 1, leaves: 1 }).unwrap(); 991 assert!(sent(&short).contains(r#""Animals":"Contains: Birds, and 1 more. For example: Owl""#), "{}", sent(&short)); 992 } 993 994 /// Contract: five 0.9 edges are a 0.59 path and lose to one 0.8 edge, 995 /// which the geometric mean (0.90 against 0.8) would get backwards. 996 #[test] 997 fn five_confident_edges_lose_to_one_likelier_edge() { 998 // root → deep (0.9) → d1 (0.9) → d2 (0.9) → d3 (0.9) → deep-leaf (0.9); 999 // root → shallow (0.8, a leaf). 1000 let deep = Branch::node( 1001 "deep", 1002 [ 1003 Branch::node( 1004 "d1", 1005 [ 1006 Branch::node( 1007 "d2", 1008 [ 1009 Branch::node("d3", [Branch::leaf("deep-leaf"), Branch::leaf("d3-x")]), 1010 Branch::leaf("d2-x"), 1011 ], 1012 ), 1013 Branch::leaf("d1-x"), 1014 ], 1015 ), 1016 Branch::leaf("deep-x"), 1017 ], 1018 ); 1019 let tree = Tree::new(Branch::node("", [deep, Branch::leaf("shallow")])).unwrap(); 1020 let (o, _) = run(&tree, plain(), |path| match path { 1021 [] => vec![0.9, 0.8], 1022 _ => vec![0.9, 0.1], 1023 }); 1024 let Outcome::Leaf { leaf, probability, separation } = o else { panic!("{o:?}") }; 1025 assert_eq!(tree.path(leaf), ["shallow"]); 1026 assert!((probability - 0.8).abs() < 1e-12, "{probability}"); 1027 // The runner-up is `deep-leaf` itself, at 0.9⁵ ≈ 0.59. 1028 assert!((0.9f64.powi(5) - 0.59).abs() < 0.001); 1029 assert!((separation.unwrap() - 0.8 / 0.9f64.powi(5)).abs() < 1e-9, "{separation:?}"); 1030 } 1031 1032 /// Exactness: the result equals exhaustive enumeration of every 1033 /// leaf, for every K, on a tree where greedy descent is wrong. 1034 #[test] 1035 fn exact_against_exhaustive() { 1036 let tree = Tree::new(Branch::node( 1037 "", 1038 [ 1039 Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2"), Branch::leaf("a3")]), 1040 Branch::node("b", [Branch::leaf("b1"), Branch::node("b2", [Branch::leaf("b2x"), Branch::leaf("b2y")])]), 1041 Branch::node("c", [Branch::leaf("c1"), Branch::leaf("c2")]), 1042 ], 1043 )) 1044 .unwrap(); 1045 // Greedy takes `a` (0.4) but its leaves split three ways; `b1` is 1046 // 0.35 × 0.95. 1047 let answer = |path: &[&str]| -> Vec<f64> { 1048 match path { 1049 [] => vec![0.4, 0.35, 0.25], 1050 ["a"] => vec![0.34, 0.33, 0.33], 1051 ["b"] => vec![0.95, 0.05], 1052 ["b", "b2"] => vec![0.5, 0.5], 1053 ["c"] => vec![0.6, 0.4], 1054 other => panic!("{other:?} asked"), 1055 } 1056 }; 1057 let mut leaves: Vec<(f64, Vec<&str>)> = Vec::new(); 1058 let mut stack = vec![(tree.root(), 0.0)]; 1059 while let Some((n, s)) = stack.pop() { 1060 if tree.is_leaf(n) { 1061 leaves.push((s, tree.path(n))); 1062 continue; 1063 } 1064 for (&c, p) in tree.children(n).iter().zip(answer(&tree.path(n))) { 1065 stack.push((c, s + p.max(PROBABILITY_FLOOR).ln())); 1066 } 1067 } 1068 leaves.sort_by(|a, b| b.0.total_cmp(&a.0)); 1069 // Lookahead and twins change the rounds, never the result. 1070 let extras = [ 1071 (None, None), 1072 (Some(Lookahead::default()), None), 1073 (Some(Lookahead { m: 1, budget: 250 }), TwinBelow::new(1.5)), 1074 (Some(Lookahead::default()), TwinBelow::new(1.5)), 1075 ]; 1076 for k in 1..=5 { 1077 for (lookahead, twin) in extras { 1078 let params = Params { k: NonZeroUsize::new(k).unwrap(), tau: None, lookahead, twin }; 1079 let (o, _) = run_rounds(&tree, params, fixed(answer)); 1080 let Outcome::Leaf { leaf, separation, .. } = o else { panic!() }; 1081 assert_eq!(tree.path(leaf), leaves[0].1, "K={k} {params:?}"); 1082 assert!((separation.unwrap() - (leaves[0].0 - leaves[1].0).exp()).abs() < 1e-12, "{params:?}"); 1083 } 1084 } 1085 } 1086 1087 /// A confident path never asks the unlikely subtrees below it. 1088 #[test] 1089 fn prunes_what_cannot_win() { 1090 let wide = |l: &str| Branch::node(l, (0..50).map(|i| chain(&format!("{l}{i}-"), 3, &format!("{l}{i}!")))); 1091 let tree = Tree::new(Branch::node("", [wide("x"), wide("y"), wide("z")])).unwrap(); 1092 let (o, asked) = run(&tree, plain(), |path| match path.len() { 1093 0 => vec![0.98, 0.01, 0.01], 1094 1 => std::iter::once(0.9).chain(std::iter::repeat_n(0.1 / 49.0, 49)).collect(), 1095 _ => vec![0.9, 0.1], 1096 }); 1097 assert_eq!(tree.path(leaf_of(&o)).last(), Some(&"x0!")); 1098 // The best path, and its runner-ups' first Choices; nowhere near 1099 // the 3 + 150 + 450 internal nodes. 1100 assert!(asked <= 12, "{asked}"); 1101 } 1102 1103 #[test] 1104 fn one_child_is_not_asked() { 1105 let tree = Tree::new(Branch::node( 1106 "", 1107 [Branch::node("only", [Branch::node("still", [Branch::leaf("p"), Branch::leaf("q")])])], 1108 )) 1109 .unwrap(); 1110 let (o, asked) = run(&tree, plain(), |path| { 1111 assert_eq!(path, ["only", "still"]); 1112 vec![0.3, 0.7] 1113 }); 1114 let Outcome::Leaf { leaf, probability, .. } = o else { panic!() }; 1115 assert_eq!(tree.path(leaf), ["only", "still", "q"]); 1116 assert!((probability - 0.7).abs() < 1e-12); 1117 assert_eq!(asked, 1); 1118 } 1119 1120 #[test] 1121 fn a_single_leaf_has_no_separation() { 1122 let tree = Tree::new(Branch::node("", [Branch::leaf("sole")])).unwrap(); 1123 let (o, asked) = run(&tree, plain(), |_| unreachable!()); 1124 assert_eq!(o, Outcome::Leaf { leaf: NodeId(1), probability: 1.0, separation: None }); 1125 assert_eq!(asked, 0); 1126 } 1127 1128 #[test] 1129 fn a_zero_is_floored_not_minus_infinity() { 1130 let tree = Tree::new(Branch::node( 1131 "", 1132 [Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2")]), Branch::leaf("b")], 1133 )) 1134 .unwrap(); 1135 let (o, _) = run(&tree, plain(), |path| match path { 1136 [] => vec![1.0, 0.0], 1137 _ => vec![0.0, 0.0], 1138 }); 1139 let Outcome::Leaf { leaf, separation, .. } = o else { panic!() }; 1140 // a1 and a2 tie at 1e-9; `b` (also 1e-9) is not ranked below −∞. 1141 assert_eq!(tree.path(leaf), ["a", "a1"]); 1142 assert_eq!(separation, Some(1.0)); 1143 } 1144 1145 /// Lookahead scores two levels per round trip on a chain, and asks 1146 /// nothing it was not going to use there. 1147 #[test] 1148 fn lookahead_halves_the_round_trips() { 1149 let tree = Tree::new(chain("n", 8, "end")).unwrap(); 1150 let answer = |_: &[&str]| vec![0.9, 0.1]; 1151 let k1 = Params { k: NonZeroUsize::MIN, ..plain() }; 1152 let (o, rounds) = run_rounds(&tree, k1, fixed(answer)); 1153 assert_eq!(tree.path(leaf_of(&o)).last(), Some(&"end")); 1154 assert_eq!(rounds.len(), 8); 1155 let look = Params { lookahead: Some(Lookahead { m: 1, budget: 1_000 }), ..k1 }; 1156 let (o2, rounds) = run_rounds(&tree, look, fixed(answer)); 1157 assert_eq!(o2, o); 1158 assert_eq!(rounds.len(), 4, "{rounds:?}"); 1159 assert!(rounds.iter().all(|r| matches!(r[..], [Ask::Node(_), Ask::Lookahead(_)])), "{rounds:?}"); 1160 } 1161 1162 /// Lookahead fills the round up to the budget, likeliest child 1163 /// (most leaves) first, and never past it. 1164 #[test] 1165 fn lookahead_stays_within_the_budget() { 1166 let tree = Tree::new(Branch::node( 1167 "", 1168 [ 1169 Branch::node("small", [Branch::leaf("s1"), Branch::leaf("s2")]), 1170 Branch::node("big", [Branch::leaf("b1"), Branch::leaf("b2"), Branch::leaf("b3")]), 1171 Branch::node("mid", [Branch::leaf("m1"), Branch::leaf("m2")]), 1172 ], 1173 )) 1174 .unwrap(); 1175 let [small, big, mid] = tree.children(tree.root()) else { panic!() }; 1176 let first = |budget| { 1177 let mut s = Search::new(&tree, Params { lookahead: Some(Lookahead { m: 3, budget }), ..plain() }); 1178 let Step::Ask(round) = s.step(|_| 100) else { panic!() }; 1179 round 1180 }; 1181 let root = Ask::Node(tree.root()); 1182 assert_eq!(first(99), [root]); 1183 assert_eq!(first(250), [root, Ask::Lookahead(*big)]); 1184 // Equal leaf counts keep the caller's order. 1185 assert_eq!(first(400), [root, Ask::Lookahead(*big), Ask::Lookahead(*small), Ask::Lookahead(*mid)]); 1186 } 1187 1188 /// A close call is asked again in reverse, and the two answers' 1189 /// log-probabilities are averaged, which undoes a bias toward the 1190 /// first-listed option. 1191 #[test] 1192 fn a_close_call_gets_an_order_twin() { 1193 let tree = Tree::new(Branch::node("", [Branch::leaf("a"), Branch::leaf("b")])).unwrap(); 1194 // Listed first gains 0.1: truly a 0.45, b 0.55. 1195 let biased = |_: &[&str], order: Order| match order { 1196 Order::Caller => vec![0.55, 0.45], 1197 Order::Reversed => vec![0.65, 0.35], 1198 }; 1199 let (o, rounds) = run_rounds(&tree, plain(), biased); 1200 assert_eq!(tree.path(leaf_of(&o)), ["a"]); 1201 assert_eq!(rounds.len(), 1); 1202 1203 let twin = Params { twin: TwinBelow::new(1.5), ..plain() }; 1204 let (o, rounds) = run_rounds(&tree, twin, biased); 1205 assert_eq!(rounds, [vec![Ask::Node(tree.root())], vec![Ask::Twin(tree.root())]]); 1206 let Outcome::Leaf { leaf, probability, separation } = o else { panic!() }; 1207 assert_eq!(tree.path(leaf), ["b"]); 1208 let (a, b) = ((0.55f64 * 0.35).sqrt(), (0.45f64 * 0.65).sqrt()); 1209 assert!((probability - b).abs() < 1e-12, "{probability}"); 1210 assert!((separation.unwrap() - b / a).abs() < 1e-12); 1211 1212 // Not close: no twin. 1213 let (_, rounds) = run_rounds(&tree, twin, fixed(|_| vec![0.2, 0.8])); 1214 assert_eq!(rounds.len(), 1); 1215 assert!(TwinBelow::new(1.0).is_none()); 1216 assert!(TwinBelow::new(f64::INFINITY).is_none()); 1217 } 1218 1219 /// A close call found by lookahead waits for its twin too. 1220 #[test] 1221 fn a_close_lookahead_answer_is_twinned() { 1222 let tree = Tree::new(Branch::node( 1223 "", 1224 [Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2")]), Branch::leaf("b")], 1225 )) 1226 .unwrap(); 1227 let a = tree.children(tree.root())[0]; 1228 let (o, rounds) = run_rounds(&tree, Params::default(), fixed(|path| match path { 1229 [] => vec![0.9, 0.1], 1230 _ => vec![0.5, 0.5], 1231 })); 1232 assert_eq!(rounds, [vec![Ask::Node(tree.root()), Ask::Lookahead(a)], vec![Ask::Twin(a)]]); 1233 assert_eq!(tree.path(leaf_of(&o)), ["a", "a1"]); 1234 } 1235 1236 fn tau(t: f64) -> Params { 1237 Params { tau: Tau::new(t), ..plain() } 1238 } 1239 1240 #[test] 1241 fn tau_stops_where_the_path_falls_below_it() { 1242 let tree = Tree::new(chain("n", 4, "end")).unwrap(); 1243 // Root is n0; each level keeps 0.8 on the chain: 0.8, 0.64, 0.51, 0.41. 1244 let answer = |_: &[&str]| vec![0.8, 0.2]; 1245 let (o, asked) = run(&tree, tau(0.5), answer); 1246 let Outcome::Stop { node, depth, leaf, probability } = o else { panic!("{o:?}") }; 1247 assert_eq!(tree.path(node), ["n1", "n2", "n3"]); 1248 assert_eq!(depth, 3); 1249 assert!(!leaf); 1250 assert!((probability - 0.512).abs() < 1e-12); 1251 // n3 is asked, since only its answer shows its children fall 1252 // below 0.5; nothing below τ is asked. 1253 assert_eq!(asked, 4); 1254 1255 let (o, _) = run(&tree, tau(0.4), answer); 1256 let Outcome::Stop { node, leaf, .. } = o else { panic!() }; 1257 assert_eq!(tree.path(node), ["n1", "n2", "n3", "end"]); 1258 assert!(leaf); 1259 } 1260 1261 #[test] 1262 fn tau_prefers_depth_then_probability() { 1263 let tree = Tree::new(Branch::node( 1264 "", 1265 [ 1266 Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2")]), 1267 Branch::node("b", [Branch::node("b1", [Branch::leaf("x"), Branch::leaf("y")]), Branch::leaf("b2")]), 1268 ], 1269 )) 1270 .unwrap(); 1271 let (o, _) = run(&tree, tau(0.2), |path| match path { 1272 [] => vec![0.6, 0.4], 1273 ["a"] => vec![0.5, 0.5], 1274 ["b"] => vec![0.6, 0.4], 1275 ["b", "b1"] => vec![0.5, 0.5], 1276 other => panic!("{other:?}"), 1277 }); 1278 // a1/a2 are 0.3 at depth 2; b1 is 0.24 at depth 2, and x/y 0.12 < τ. 1279 // So depth 2, likeliest: a1. 1280 let Outcome::Stop { node, depth, .. } = o else { panic!() }; 1281 assert_eq!((tree.path(node), depth), (vec!["a", "a1"], 2)); 1282 } 1283 1284 #[test] 1285 fn tau_is_in_zero_to_one() { 1286 assert!(Tau::new(0.0).is_none()); 1287 assert!(Tau::new(1.5).is_none()); 1288 assert!(Tau::new(f64::NAN).is_none()); 1289 assert!(Tau::new(1.0).is_some()); 1290 } 1291 1292 #[test] 1293 fn trees_that_cannot_be_asked_are_refused() { 1294 assert_eq!(Tree::new(Branch::leaf("")).unwrap_err(), TreeError::NoChoice); 1295 let wide = Branch::node("", (0..256).map(|i| Branch::leaf(i.to_string()))); 1296 assert!(matches!(Tree::new(wide), Err(TreeError::Fanout { children: 256, .. }))); 1297 assert!(Tree::new(Branch::node("", (0..255).map(|i| Branch::leaf(i.to_string())))).is_ok()); 1298 let dup = Branch::node("", [Branch::node("a", [Branch::leaf("x"), Branch::leaf("x")])]); 1299 assert_eq!( 1300 Tree::new(dup).unwrap_err(), 1301 TreeError::DuplicateLabel { path: vec!["a".into()], label: "x".into() } 1302 ); 1303 let empty = Branch::node("", [Branch::leaf("")]); 1304 assert_eq!(Tree::new(empty).unwrap_err(), TreeError::EmptyLabel { path: vec![] }); 1305 } 1306 1307 #[test] 1308 fn answers_are_checked() { 1309 let tree = Tree::new(Branch::node("", [Branch::leaf("a"), Branch::leaf("b")])).unwrap(); 1310 let mut s = Search::new(&tree, Params::default()); 1311 let root = Ask::Node(tree.root()); 1312 let Step::Ask(round) = s.step(|_| 100) else { panic!() }; 1313 assert_eq!(round, [root]); 1314 assert_eq!(s.step(|_| 100), Step::Ask(round.clone())); 1315 assert_eq!(s.answer(Ask::Node(NodeId(1)), &[1.0]), Err(AnswerError::NotAsked(Ask::Node(NodeId(1))))); 1316 assert_eq!(s.answer(Ask::Twin(tree.root()), &[0.5, 0.5]), Err(AnswerError::NotAsked(Ask::Twin(tree.root())))); 1317 assert!(matches!(s.answer(root, &[1.0]), Err(AnswerError::Count { expected: 2, got: 1, .. }))); 1318 assert!(matches!(s.answer(root, &[0.5, f64::NAN]), Err(AnswerError::NotAProbability { .. }))); 1319 assert!(matches!(s.answer(root, &[0.5, 1.2]), Err(AnswerError::NotAProbability { .. }))); 1320 s.answer(root, &[0.25, 0.75]).unwrap(); 1321 let Step::Done(o) = s.step(|_| 100) else { panic!() }; 1322 assert_eq!(tree.path(leaf_of(&o)), ["b"]); 1323 } 1324}