use super::drafter::{validate_candidates, Drafter, TreeContextView};
use anyhow::{anyhow, ensure, Result};
use std::cmp::Ordering;
use std::collections::BinaryHeap;
#[derive(Debug, Clone, Copy)]
pub struct DynamicTreeConfig {
pub budget: usize,
pub max_depth: usize,
pub top_k: usize,
}
impl Default for DynamicTreeConfig {
fn default() -> Self {
Self {
budget: 64,
max_depth: 8,
top_k: 10,
}
}
}
impl DynamicTreeConfig {
pub fn validate(&self) -> Result<()> {
ensure!(self.budget >= 1, "budget must be >= 1");
ensure!(self.max_depth >= 1, "max_depth must be >= 1");
ensure!(self.top_k >= 1, "top_k must be >= 1");
ensure!(
self.budget <= 8192,
"budget {} exceeds sane upper bound 8192",
self.budget
);
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct ExpandedTree {
pub tokens: Vec<u32>,
pub parents: Vec<Option<usize>>,
pub depths: Vec<usize>,
pub cum_log_probs: Vec<f64>,
}
impl ExpandedTree {
pub fn len(&self) -> usize {
self.tokens.len()
}
pub fn is_empty(&self) -> bool {
self.tokens.is_empty()
}
pub fn build_tree_mask(&self, prefix_len: usize) -> Result<Vec<f32>> {
const ATTENDED: f32 = 0.0;
const MASKED: f32 = -65504.0;
let q = self.len();
let mask_stride = prefix_len.checked_add(q).ok_or_else(|| {
anyhow!(
"build_tree_mask: prefix_len ({}) + tree.len ({}) overflows usize",
prefix_len,
q
)
})?;
let total = q.checked_mul(mask_stride).ok_or_else(|| {
anyhow!(
"build_tree_mask: tree.len ({}) * mask_stride ({}) overflows usize",
q,
mask_stride
)
})?;
let mut mask = vec![MASKED; total];
for iq1 in 0..q {
let row_base = iq1 * mask_stride;
for k in 0..prefix_len {
mask[row_base + k] = ATTENDED;
}
let mut cur = Some(iq1);
while let Some(node) = cur {
mask[row_base + prefix_len + node] = ATTENDED;
cur = self.parents[node];
}
}
Ok(mask)
}
pub fn validate(&self) -> Result<()> {
let n = self.len();
ensure!(n >= 1, "tree is empty");
ensure!(
self.parents.len() == n && self.depths.len() == n && self.cum_log_probs.len() == n,
"tree vec lengths inconsistent"
);
ensure!(
self.parents[0].is_none(),
"root (index 0) must have no parent"
);
ensure!(self.depths[0] == 0, "root depth must be 0");
ensure!(
self.cum_log_probs[0] == 0.0_f64,
"root cum_log_prob must be exactly 0.0"
);
for i in 1..n {
let p =
self.parents[i].ok_or_else(|| anyhow!("non-root node {} must have a parent", i))?;
ensure!(
p < i,
"parents[{}] = {} violates topological order (parent must precede child)",
i,
p
);
ensure!(
self.depths[i] == self.depths[p] + 1,
"depths[{}] = {} but depths[parent={}] + 1 = {}",
i,
self.depths[i],
p,
self.depths[p] + 1
);
ensure!(
self.cum_log_probs[i].is_finite(),
"cum_log_probs[{}] is not finite",
i
);
}
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
struct PendingCandidate {
parent_idx: usize,
token: u32,
cum_log_prob: f64,
seq: usize,
}
impl PartialEq for PendingCandidate {
fn eq(&self, other: &Self) -> bool {
self.cum_log_prob == other.cum_log_prob && self.seq == other.seq
}
}
impl Eq for PendingCandidate {}
impl Ord for PendingCandidate {
fn cmp(&self, other: &Self) -> Ordering {
self.cum_log_prob
.partial_cmp(&other.cum_log_prob)
.unwrap_or(Ordering::Equal)
.then_with(|| other.seq.cmp(&self.seq))
}
}
impl PartialOrd for PendingCandidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
pub fn expand_dynamic_tree<D: Drafter>(
root_token: u32,
drafter: &mut D,
cfg: &DynamicTreeConfig,
) -> Result<ExpandedTree> {
cfg.validate()?;
let mut tokens: Vec<u32> = Vec::with_capacity(cfg.budget);
let mut parents: Vec<Option<usize>> = Vec::with_capacity(cfg.budget);
let mut depths: Vec<usize> = Vec::with_capacity(cfg.budget);
let mut cum: Vec<f64> = Vec::with_capacity(cfg.budget);
tokens.push(root_token);
parents.push(None);
depths.push(0);
cum.push(0.0);
let mut heap = BinaryHeap::<PendingCandidate>::new();
let mut seq_counter: usize = 0;
if cfg.budget > 1 && cfg.max_depth >= 1 {
let view = TreeContextView {
tokens: &tokens,
parents: &parents,
};
let candidates = drafter.predict_topk(view, 0, cfg.top_k)?;
validate_candidates(&candidates, cfg.top_k)?;
for cand in candidates {
let child_cum = cand.log_prob as f64;
ensure!(
child_cum.is_finite(),
"seed cum_log_prob not finite: {}",
child_cum
);
heap.push(PendingCandidate {
parent_idx: 0,
token: cand.token,
cum_log_prob: child_cum,
seq: seq_counter,
});
seq_counter += 1;
}
}
while tokens.len() < cfg.budget {
let Some(pending) = heap.pop() else {
break; };
let parent_idx = pending.parent_idx;
let child_idx = tokens.len();
let new_depth = depths[parent_idx] + 1;
tokens.push(pending.token);
parents.push(Some(parent_idx));
depths.push(new_depth);
cum.push(pending.cum_log_prob);
if tokens.len() < cfg.budget && new_depth < cfg.max_depth {
let view = TreeContextView {
tokens: &tokens,
parents: &parents,
};
let candidates = drafter.predict_topk(view, child_idx, cfg.top_k)?;
validate_candidates(&candidates, cfg.top_k)?;
let parent_cum = cum[child_idx];
for cand in candidates {
let child_cum = parent_cum + (cand.log_prob as f64);
ensure!(
child_cum.is_finite(),
"cumulative log_prob overflowed at depth {} (parent_cum={}, edge={})",
new_depth + 1,
parent_cum,
cand.log_prob
);
heap.push(PendingCandidate {
parent_idx: child_idx,
token: cand.token,
cum_log_prob: child_cum,
seq: seq_counter,
});
seq_counter += 1;
}
}
}
let out = ExpandedTree {
tokens,
parents,
depths,
cum_log_probs: cum,
};
out.validate()
.map_err(|e| anyhow!("expand_dynamic_tree produced invalid tree: {}", e))?;
Ok(out)
}
pub trait CacheControlDrafter: Drafter {
fn cache_len(&self) -> usize;
fn clear_cache(&mut self);
}
pub fn expand_dynamic_tree_with_cache<D: CacheControlDrafter>(
root_token: u32,
drafter: &mut D,
cfg: &DynamicTreeConfig,
) -> Result<ExpandedTree> {
cfg.validate()?;
drafter.clear_cache();
let mut tokens: Vec<u32> = Vec::with_capacity(cfg.budget);
let mut parents: Vec<Option<usize>> = Vec::with_capacity(cfg.budget);
let mut depths: Vec<usize> = Vec::with_capacity(cfg.budget);
let mut cum: Vec<f64> = Vec::with_capacity(cfg.budget);
tokens.push(root_token);
parents.push(None);
depths.push(0);
cum.push(0.0);
let mut heap = BinaryHeap::<PendingCandidate>::new();
let mut seq_counter: usize = 0;
if cfg.budget > 1 && cfg.max_depth >= 1 {
let view = TreeContextView {
tokens: &tokens,
parents: &parents,
};
let candidates = drafter.predict_topk(view, 0, cfg.top_k)?;
validate_candidates(&candidates, cfg.top_k)?;
debug_assert_eq!(drafter.cache_len(), 1);
for cand in candidates {
let child_cum = cand.log_prob as f64;
ensure!(
child_cum.is_finite(),
"seed cum_log_prob not finite: {}",
child_cum
);
heap.push(PendingCandidate {
parent_idx: 0,
token: cand.token,
cum_log_prob: child_cum,
seq: seq_counter,
});
seq_counter += 1;
}
}
while tokens.len() < cfg.budget {
let Some(pending) = heap.pop() else {
break;
};
let parent_idx = pending.parent_idx;
let child_idx = tokens.len();
let new_depth = depths[parent_idx] + 1;
tokens.push(pending.token);
parents.push(Some(parent_idx));
depths.push(new_depth);
cum.push(pending.cum_log_prob);
if tokens.len() < cfg.budget && new_depth < cfg.max_depth {
let view = TreeContextView {
tokens: &tokens,
parents: &parents,
};
let candidates = drafter.predict_topk(view, child_idx, cfg.top_k)?;
validate_candidates(&candidates, cfg.top_k)?;
let parent_cum = cum[child_idx];
for cand in candidates {
let child_cum = parent_cum + (cand.log_prob as f64);
ensure!(
child_cum.is_finite(),
"cumulative log_prob overflowed at depth {} (parent_cum={}, edge={})",
new_depth + 1,
parent_cum,
cand.log_prob
);
heap.push(PendingCandidate {
parent_idx: child_idx,
token: cand.token,
cum_log_prob: child_cum,
seq: seq_counter,
});
seq_counter += 1;
}
}
}
let out = ExpandedTree {
tokens,
parents,
depths,
cum_log_probs: cum,
};
out.validate().map_err(|e| {
anyhow!(
"expand_dynamic_tree_with_cache produced invalid tree: {}",
e
)
})?;
Ok(out)
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::super::drafter::{BiasedMockDrafter, DraftCandidate, MockDrafter};
use super::*;
use std::collections::HashSet;
struct ScriptedDrafter {
scripts: Vec<Vec<DraftCandidate>>,
call_count: usize,
}
impl Drafter for ScriptedDrafter {
fn predict_topk(
&mut self,
_tree: super::super::drafter::TreeContextView<'_>,
_node_to_expand: usize,
_top_k: usize,
) -> Result<Vec<DraftCandidate>> {
let script = self.scripts[self.call_count].clone();
self.call_count += 1;
Ok(script)
}
}
struct CacheTrackingMock<D: Drafter> {
inner: D,
cache_slots: Vec<u64>,
next_appended_tag: u64,
rollback_history: Vec<Vec<usize>>,
predict_entry_lens: Vec<usize>,
}
impl<D: Drafter> CacheTrackingMock<D> {
fn new(inner: D) -> Self {
Self {
inner,
cache_slots: Vec::new(),
next_appended_tag: 1,
rollback_history: Vec::new(),
predict_entry_lens: Vec::new(),
}
}
}
impl<D: Drafter> Drafter for CacheTrackingMock<D> {
fn predict_topk(
&mut self,
tree: super::super::drafter::TreeContextView<'_>,
node_to_expand: usize,
top_k: usize,
) -> Result<Vec<DraftCandidate>> {
self.predict_entry_lens.push(self.cache_slots.len());
let mut cursor = tree.parents[node_to_expand];
let mut required_ancestor_count = 0;
while let Some(idx) = cursor {
required_ancestor_count += 1;
assert!(
required_ancestor_count <= self.cache_slots.len(),
"expanding node {} but cache.len()={} < ancestor count {} \
(ancestor not yet expanded — orchestrator bug)",
node_to_expand,
self.cache_slots.len(),
required_ancestor_count,
);
cursor = tree.parents[idx];
}
let candidates = self.inner.predict_topk(tree, node_to_expand, top_k)?;
self.cache_slots.push(self.next_appended_tag);
self.next_appended_tag += 1;
Ok(candidates)
}
}
impl<D: Drafter> CacheControlDrafter for CacheTrackingMock<D> {
fn cache_len(&self) -> usize {
self.cache_slots.len()
}
fn clear_cache(&mut self) {
self.cache_slots.clear();
}
}
#[test]
fn adr_037_e6_cache_orchestrator_root_only_no_expansion_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 1,
max_depth: 4,
top_k: 3,
};
let inner = MockDrafter::default();
let mut mock = CacheTrackingMock::new(inner);
let tree = expand_dynamic_tree_with_cache(123, &mut mock, &cfg).expect("expand");
assert_eq!(tree.len(), 1);
assert_eq!(mock.cache_len(), 0);
assert!(mock.rollback_history.is_empty());
assert!(mock.predict_entry_lens.is_empty());
}
#[test]
fn adr_037_e6_cache_orchestrator_linear_chain_no_rollbacks_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 5,
max_depth: 4,
top_k: 1,
};
let inner = MockDrafter {
vocab_size: 1000,
base_log_prob: -0.5,
log_prob_slope: 0.0,
};
let mut mock = CacheTrackingMock::new(inner);
let tree = expand_dynamic_tree_with_cache(123, &mut mock, &cfg).expect("expand");
assert_eq!(tree.len(), 5);
assert_eq!(mock.cache_len(), 4);
assert!(
mock.rollback_history.is_empty(),
"tree-mask design: no rollbacks during expansion"
);
assert_eq!(mock.predict_entry_lens, vec![0, 1, 2, 3]);
}
#[test]
fn adr_037_e6_cache_orchestrator_sibling_expansion_no_rollback_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 6,
max_depth: 2,
top_k: 3,
};
let inner = MockDrafter {
vocab_size: 1000,
base_log_prob: -0.5,
log_prob_slope: -1.0,
};
let mut mock = CacheTrackingMock::new(inner);
let tree = expand_dynamic_tree_with_cache(0, &mut mock, &cfg).expect("expand");
assert_eq!(tree.len(), 6);
assert!(
mock.rollback_history.is_empty(),
"tree-mask design: cache never rolls back during expansion"
);
assert!(!mock.predict_entry_lens.is_empty());
}
#[test]
fn adr_037_e6_cache_orchestrator_deep_tree_cross_branch_no_panic_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 12,
max_depth: 4,
top_k: 3,
};
let inner = MockDrafter {
vocab_size: 1000,
base_log_prob: -0.5,
log_prob_slope: -1.0,
};
let mut mock = CacheTrackingMock::new(inner);
let tree = expand_dynamic_tree_with_cache(0, &mut mock, &cfg)
.expect("max_depth=4 + cross-branch must not panic");
assert!(tree.len() <= cfg.budget);
assert!(mock.cache_len() >= 1);
assert_eq!(mock.cache_len(), mock.predict_entry_lens.len());
}
#[test]
fn adr_037_e6_cache_orchestrator_clear_called_at_entry_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 2,
max_depth: 1,
top_k: 1,
};
let inner = MockDrafter::default();
let mut mock = CacheTrackingMock::new(inner);
mock.cache_slots = vec![99, 88, 77]; let _ = expand_dynamic_tree_with_cache(0, &mut mock, &cfg)
.expect("expand should clear cache at entry");
assert_eq!(mock.cache_len(), 1);
assert_eq!(mock.predict_entry_lens[0], 0);
}
#[test]
fn adr_037_e6_cache_orchestrator_no_rollback_history_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 4,
max_depth: 2,
top_k: 3,
};
let inner = MockDrafter {
vocab_size: 1000,
base_log_prob: -0.5,
log_prob_slope: -1.0,
};
let mut mock = CacheTrackingMock::new(inner);
let _ = expand_dynamic_tree_with_cache(0, &mut mock, &cfg).expect("expand");
assert!(
mock.rollback_history.is_empty(),
"orchestrator should not call rollback in tree-mask design"
);
}
#[test]
fn adr_037_e4a_budget_1_returns_only_root_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 1,
max_depth: 4,
top_k: 4,
};
let mut d = MockDrafter::default();
let tree = expand_dynamic_tree(12345, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 1);
assert_eq!(tree.tokens[0], 12345);
assert_eq!(tree.parents[0], None);
assert_eq!(tree.depths[0], 0);
assert_eq!(tree.cum_log_probs[0], 0.0);
}
#[test]
fn adr_037_e4a_linear_chain_via_top_k_1_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 5,
max_depth: 10,
top_k: 1,
};
let mut d = MockDrafter::default();
let tree = expand_dynamic_tree(100, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 5);
for i in 1..tree.len() {
assert_eq!(tree.parents[i], Some(i - 1));
assert_eq!(tree.depths[i], i);
}
}
#[test]
fn adr_037_e4a_fixed_square_via_max_depth_1_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 5, max_depth: 1, top_k: 4,
};
let mut d = MockDrafter::default();
let tree = expand_dynamic_tree(0, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 5);
for i in 1..5 {
assert_eq!(tree.parents[i], Some(0));
assert_eq!(tree.depths[i], 1);
}
}
#[test]
fn adr_037_e4a_dynamic_asymmetric_expands_biased_subtree_2026_05_22() {
let mut bias = HashSet::new();
bias.insert(1); bias.insert(3); bias.insert(5);
bias.insert(7);
let mut d = BiasedMockDrafter {
vocab_size: 32_000,
base_log_prob: -1.0,
log_prob_slope: -1.0,
bias_nodes: bias,
};
let cfg = DynamicTreeConfig {
budget: 10,
max_depth: 6,
top_k: 2,
};
let tree = expand_dynamic_tree(0, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 10);
let max_depth = *tree.depths.iter().max().unwrap();
assert!(
max_depth >= 3,
"expected biased subtree to expand to depth >= 3, got max_depth = {max_depth}"
);
}
#[test]
fn adr_037_e4a_max_depth_caps_subtree_growth_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 100,
max_depth: 2,
top_k: 3,
};
let mut d = MockDrafter::default();
let tree = expand_dynamic_tree(0, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 13);
assert_eq!(*tree.depths.iter().max().unwrap(), 2);
}
#[test]
fn adr_037_e4a_budget_exhaustion_stops_expansion_2026_05_22() {
let cfg = DynamicTreeConfig {
budget: 7,
max_depth: 10,
top_k: 3,
};
let mut d = MockDrafter::default();
let tree = expand_dynamic_tree(0, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 7);
}
#[test]
fn adr_037_e4a_global_best_first_admits_grandchild_before_sibling_2026_05_22() {
let d = ScriptedDrafter {
scripts: vec![
vec![
DraftCandidate {
token: 100,
log_prob: -0.1,
}, DraftCandidate {
token: 200,
log_prob: -2.0,
}, ],
vec![
DraftCandidate {
token: 110,
log_prob: -0.5,
}, DraftCandidate {
token: 120,
log_prob: -1.0,
}, ],
vec![
DraftCandidate {
token: 111,
log_prob: -0.5,
}, ],
],
call_count: 0,
};
let cfg = DynamicTreeConfig {
budget: 4,
max_depth: 3,
top_k: 2,
};
let mut d = d;
let tree = expand_dynamic_tree(1, &mut d, &cfg).unwrap();
assert_eq!(tree.len(), 4);
assert_eq!(tree.tokens, vec![1, 100, 110, 120]);
assert_eq!(tree.parents, vec![None, Some(0), Some(1), Some(1)]);
assert_eq!(tree.depths, vec![0, 1, 2, 2]);
let expected = [0.0, -0.1, -0.6, -1.1];
for (i, &exp) in expected.iter().enumerate() {
assert!(
(tree.cum_log_probs[i] - exp).abs() < 1e-6,
"cum[{}] = {} != {}",
i,
tree.cum_log_probs[i],
exp
);
}
assert!(
!tree.tokens.contains(&200),
"B (low-prob sibling) should NOT have been admitted ahead of A's grandchildren"
);
}
#[test]
fn adr_037_e4a_build_tree_mask_matches_phase_e1_contract_2026_05_22() {
let tree = ExpandedTree {
tokens: vec![1, 10, 20, 11],
parents: vec![None, Some(0), Some(0), Some(1)],
depths: vec![0, 1, 1, 2],
cum_log_probs: vec![0.0, -0.2, -0.5, -0.5],
};
tree.validate().expect("hand-built tree must validate");
let prefix_len = 5;
let mask = tree
.build_tree_mask(prefix_len)
.expect("build_tree_mask ok");
let q = tree.len();
let mask_stride = prefix_len + q;
assert_eq!(mask.len(), q * mask_stride);
const ATTENDED: f32 = 0.0;
const MASKED: f32 = -65504.0;
let check = |row: usize, col: usize, exp: f32, label: &str| {
assert_eq!(
mask[row * mask_stride + col],
exp,
"row {row} col {col} ({label})"
);
};
for k in 0..prefix_len {
check(0, k, ATTENDED, "prefix");
}
check(0, 5, ATTENDED, "self");
check(0, 6, MASKED, "child0 sibling");
check(0, 7, MASKED, "child1 sibling");
check(0, 8, MASKED, "grandchild");
for k in 0..prefix_len {
check(3, k, ATTENDED, "prefix");
}
check(3, 5, ATTENDED, "root ancestor");
check(3, 6, ATTENDED, "child0 parent");
check(3, 7, MASKED, "child1 sibling — not ancestor");
check(3, 8, ATTENDED, "self");
}
#[test]
fn adr_037_e4a_validate_invalid_config_rejected_2026_05_22() {
let mut d = MockDrafter::default();
let mut cfg = DynamicTreeConfig::default();
cfg.budget = 0;
assert!(expand_dynamic_tree(0, &mut d, &cfg).is_err());
let mut cfg = DynamicTreeConfig::default();
cfg.top_k = 0;
assert!(expand_dynamic_tree(0, &mut d, &cfg).is_err());
let mut cfg = DynamicTreeConfig::default();
cfg.max_depth = 0;
assert!(expand_dynamic_tree(0, &mut d, &cfg).is_err());
}
#[test]
fn adr_037_e4a_drafter_returning_unsorted_is_rejected_2026_05_22() {
struct BadDrafter;
impl Drafter for BadDrafter {
fn predict_topk(
&mut self,
_tree: super::super::drafter::TreeContextView<'_>,
_node_to_expand: usize,
_top_k: usize,
) -> Result<Vec<DraftCandidate>> {
Ok(vec![
DraftCandidate {
token: 10,
log_prob: -1.0,
},
DraftCandidate {
token: 11,
log_prob: -0.5,
},
])
}
}
let cfg = DynamicTreeConfig::default();
let mut d = BadDrafter;
let err = expand_dynamic_tree(0, &mut d, &cfg)
.unwrap_err()
.to_string();
assert!(err.contains("sorted descending"), "got: {err}");
}
#[test]
fn adr_037_e4a_expanded_tree_validate_catches_corruption_2026_05_22() {
let bad = ExpandedTree {
tokens: vec![1, 2, 3],
parents: vec![None, Some(2), Some(0)], depths: vec![0, 1, 1],
cum_log_probs: vec![0.0_f64, -0.5, -0.5],
};
let err = bad.validate().unwrap_err().to_string();
assert!(err.contains("topological"), "got: {err}");
}
}