use anyhow::{ensure, Result};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct DraftCandidate {
pub token: u32,
pub log_prob: f32,
}
pub const LOG_PROB_FLOOR: f32 = f32::MIN / 2.0;
#[derive(Debug, Clone, Copy)]
pub struct TreeContextView<'a> {
pub tokens: &'a [u32],
pub parents: &'a [Option<usize>],
}
impl<'a> TreeContextView<'a> {
pub fn path_tokens(&self, node_idx: usize) -> Vec<u32> {
let mut rev: Vec<u32> = Vec::new();
let mut cur = Some(node_idx);
while let Some(i) = cur {
rev.push(self.tokens[i]);
cur = self.parents[i];
}
rev.reverse();
rev
}
}
pub trait Drafter {
fn predict_topk(
&mut self,
tree: TreeContextView<'_>,
node_to_expand: usize,
top_k: usize,
) -> Result<Vec<DraftCandidate>>;
}
pub fn extract_top_k_from_row_logits(
row_logits: &[f32],
top_k: usize,
) -> Result<Vec<DraftCandidate>> {
ensure!(!row_logits.is_empty(), "extract_top_k: row_logits is empty");
ensure!(top_k > 0, "extract_top_k: top_k must be > 0");
for (i, &v) in row_logits.iter().enumerate() {
ensure!(
v.is_finite(),
"extract_top_k: row_logits[{}] = {} is not finite",
i,
v
);
}
ensure!(
row_logits.len() <= (u32::MAX as usize),
"extract_top_k: row_logits.len() ({}) exceeds u32::MAX",
row_logits.len()
);
let effective_k = top_k.min(row_logits.len());
let max_logit: f32 = row_logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum_exp: f64 = 0.0;
for &v in row_logits.iter() {
sum_exp += ((v - max_logit) as f64).exp();
}
let log_sumexp = (max_logit as f64) + sum_exp.ln();
use std::cmp::Reverse;
use std::collections::BinaryHeap;
#[derive(Debug, Clone, Copy)]
struct LogProbTokenF64 {
log_prob: f64,
token: u32,
}
impl PartialEq for LogProbTokenF64 {
fn eq(&self, other: &Self) -> bool {
self.log_prob == other.log_prob && self.token == other.token
}
}
impl Eq for LogProbTokenF64 {}
impl Ord for LogProbTokenF64 {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.log_prob
.partial_cmp(&other.log_prob)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| other.token.cmp(&self.token))
}
}
impl PartialOrd for LogProbTokenF64 {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
let mut heap: BinaryHeap<Reverse<LogProbTokenF64>> = BinaryHeap::with_capacity(effective_k + 1);
for (i, &v) in row_logits.iter().enumerate() {
let log_prob = (v as f64) - log_sumexp;
let candidate = LogProbTokenF64 {
log_prob,
token: i as u32,
};
if heap.len() < effective_k {
heap.push(Reverse(candidate));
} else if let Some(Reverse(min)) = heap.peek() {
if candidate > *min {
heap.pop();
heap.push(Reverse(candidate));
}
}
}
let mut out: Vec<DraftCandidate> = heap
.into_iter()
.map(|Reverse(lpt)| {
let lp = if lpt.log_prob.is_finite() {
(lpt.log_prob as f32).max(LOG_PROB_FLOOR)
} else {
LOG_PROB_FLOOR
};
DraftCandidate {
token: lpt.token,
log_prob: lp,
}
})
.collect();
out.sort_by(|a, b| {
b.log_prob
.partial_cmp(&a.log_prob)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.token.cmp(&b.token))
});
Ok(out)
}
pub fn validate_candidates(candidates: &[DraftCandidate], top_k: usize) -> Result<()> {
ensure!(
candidates.len() <= top_k,
"drafter returned {} candidates, exceeds top_k {}",
candidates.len(),
top_k
);
for (i, c) in candidates.iter().enumerate() {
ensure!(
c.log_prob.is_finite(),
"candidate[{}].log_prob = {} is not finite",
i,
c.log_prob
);
}
let mut seen = std::collections::HashSet::with_capacity(candidates.len());
for (i, c) in candidates.iter().enumerate() {
ensure!(
seen.insert(c.token),
"candidate[{}].token = {} duplicated in top-K",
i,
c.token
);
}
for w in candidates.windows(2) {
ensure!(
w[0].log_prob >= w[1].log_prob,
"candidates must be sorted descending by log_prob, got {} then {}",
w[0].log_prob,
w[1].log_prob
);
}
Ok(())
}
#[derive(Debug, Clone)]
pub struct MockDrafter {
pub vocab_size: u32,
pub base_log_prob: f32,
pub log_prob_slope: f32,
}
impl Default for MockDrafter {
fn default() -> Self {
Self {
vocab_size: 32_000,
base_log_prob: -0.5,
log_prob_slope: -0.5,
}
}
}
impl Drafter for MockDrafter {
fn predict_topk(
&mut self,
_tree: TreeContextView<'_>,
node_to_expand: usize,
top_k: usize,
) -> Result<Vec<DraftCandidate>> {
let mut out = Vec::with_capacity(top_k);
for j in 0..top_k {
let token = ((node_to_expand * 1000 + j) as u32) % self.vocab_size;
let log_prob = self.base_log_prob + (j as f32) * self.log_prob_slope;
out.push(DraftCandidate { token, log_prob });
}
Ok(out)
}
}
#[derive(Debug, Clone)]
pub struct BiasedMockDrafter {
pub vocab_size: u32,
pub base_log_prob: f32,
pub log_prob_slope: f32,
pub bias_nodes: std::collections::HashSet<usize>,
}
impl Drafter for BiasedMockDrafter {
fn predict_topk(
&mut self,
_tree: TreeContextView<'_>,
node_to_expand: usize,
top_k: usize,
) -> Result<Vec<DraftCandidate>> {
let mut out = Vec::with_capacity(top_k);
let is_biased = self.bias_nodes.contains(&node_to_expand);
for j in 0..top_k {
let token = ((node_to_expand * 1000 + j + 7) as u32) % self.vocab_size;
let log_prob = if is_biased {
-0.1 + (j as f32) * (-0.05)
} else {
self.base_log_prob + (j as f32) * self.log_prob_slope
};
out.push(DraftCandidate { token, log_prob });
}
Ok(out)
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn adr_037_e4a_validate_rejects_nan_log_prob_2026_05_22() {
let bad = vec![DraftCandidate {
token: 10,
log_prob: f32::NAN,
}];
let err = validate_candidates(&bad, 1).unwrap_err().to_string();
assert!(err.contains("not finite"), "got: {err}");
}
#[test]
fn adr_037_e4a_validate_rejects_inf_log_prob_2026_05_22() {
let bad = vec![DraftCandidate {
token: 10,
log_prob: f32::NEG_INFINITY,
}];
let err = validate_candidates(&bad, 1).unwrap_err().to_string();
assert!(err.contains("not finite"), "got: {err}");
}
#[test]
fn adr_037_e4a_validate_rejects_duplicate_tokens_2026_05_22() {
let bad = vec![
DraftCandidate {
token: 10,
log_prob: -0.5,
},
DraftCandidate {
token: 10,
log_prob: -1.0,
},
];
let err = validate_candidates(&bad, 2).unwrap_err().to_string();
assert!(err.contains("duplicated in top-K"), "got: {err}");
}
#[test]
fn adr_037_e4a_validate_rejects_unsorted_2026_05_22() {
let bad = vec![
DraftCandidate {
token: 10,
log_prob: -1.0,
},
DraftCandidate {
token: 11,
log_prob: -0.5,
}, ];
let err = validate_candidates(&bad, 2).unwrap_err().to_string();
assert!(err.contains("sorted descending"), "got: {err}");
}
#[test]
fn adr_037_e4a_validate_rejects_too_many_candidates_2026_05_22() {
let bad = vec![
DraftCandidate {
token: 10,
log_prob: -0.5,
},
DraftCandidate {
token: 11,
log_prob: -1.0,
},
DraftCandidate {
token: 12,
log_prob: -1.5,
},
];
let err = validate_candidates(&bad, 2).unwrap_err().to_string();
assert!(err.contains("exceeds top_k"), "got: {err}");
}
fn empty_tree_view() -> TreeContextView<'static> {
static EMPTY_TOKENS: &[u32] = &[];
static EMPTY_PARENTS: &[Option<usize>] = &[];
TreeContextView {
tokens: EMPTY_TOKENS,
parents: EMPTY_PARENTS,
}
}
#[test]
fn adr_037_e4a_mock_drafter_produces_unique_descending_candidates_2026_05_22() {
let mut d = MockDrafter::default();
let cands = d.predict_topk(empty_tree_view(), 0, 4).unwrap();
validate_candidates(&cands, 4).expect("mock must be valid");
assert_eq!(cands.len(), 4);
}
#[test]
fn adr_037_e4a_mock_drafter_uses_parent_idx_in_token_2026_05_22() {
let mut d = MockDrafter::default();
let c0 = d.predict_topk(empty_tree_view(), 0, 1).unwrap();
let c5 = d.predict_topk(empty_tree_view(), 5, 1).unwrap();
assert_ne!(c0[0].token, c5[0].token, "different node → different token");
}
#[test]
fn adr_037_e4a_biased_drafter_shallower_slope_at_bias_nodes_2026_05_22() {
let mut bias_nodes = std::collections::HashSet::new();
bias_nodes.insert(3);
let mut d = BiasedMockDrafter {
vocab_size: 1000,
base_log_prob: -0.5,
log_prob_slope: -0.5,
bias_nodes,
};
let unbiased = d.predict_topk(empty_tree_view(), 0, 2).unwrap();
let biased = d.predict_topk(empty_tree_view(), 3, 2).unwrap();
assert!(biased[0].log_prob > unbiased[0].log_prob);
validate_candidates(&biased, 2).unwrap();
}
#[test]
fn adr_037_e4a_tree_context_view_path_tokens_walks_to_root_2026_05_22() {
let tokens = [10, 20, 30];
let parents: [Option<usize>; 3] = [None, Some(0), Some(1)];
let view = TreeContextView {
tokens: &tokens,
parents: &parents,
};
assert_eq!(view.path_tokens(0), vec![10]);
assert_eq!(view.path_tokens(1), vec![10, 20]);
assert_eq!(view.path_tokens(2), vec![10, 20, 30]);
}
#[test]
fn adr_037_e4b10a_top_k_basic_descending_2026_05_22() {
let logits = vec![3.0f32, 1.0, 2.0, 0.0];
let out = extract_top_k_from_row_logits(&logits, 3).unwrap();
assert_eq!(out.len(), 3);
assert_eq!(out[0].token, 0);
assert_eq!(out[1].token, 2);
assert_eq!(out[2].token, 1);
validate_candidates(&out, 3).expect("must pass validate_candidates");
}
#[test]
fn adr_037_e4b10a_top_k_log_probs_sum_via_softmax_2026_05_22() {
let logits = vec![0.0f32; 4];
let out = extract_top_k_from_row_logits(&logits, 4).unwrap();
assert_eq!(out.len(), 4);
let expected = -(4.0f32).ln();
for c in &out {
assert!(
(c.log_prob - expected).abs() < 1e-5,
"log_prob {} != {expected}",
c.log_prob
);
}
}
#[test]
fn adr_037_e4b10a_top_k_returns_at_most_vocab_size_2026_05_22() {
let logits = vec![1.0f32, 2.0, 3.0];
let out = extract_top_k_from_row_logits(&logits, 10).unwrap();
assert_eq!(out.len(), 3, "vocab=3 caps the output count");
assert_eq!(out[0].token, 2);
assert_eq!(out[1].token, 1);
assert_eq!(out[2].token, 0);
}
#[test]
fn adr_037_e4b10a_top_k_rejects_empty_logits_2026_05_22() {
let err = extract_top_k_from_row_logits(&[], 1).unwrap_err();
assert!(err.to_string().contains("empty"), "got: {err}");
}
#[test]
fn adr_037_e4b10a_top_k_rejects_top_k_zero_2026_05_22() {
let err = extract_top_k_from_row_logits(&[1.0, 2.0], 0).unwrap_err();
assert!(err.to_string().contains("top_k must be > 0"), "got: {err}");
}
#[test]
fn adr_037_e4b10a_top_k_rejects_nan_logit_2026_05_22() {
let logits = vec![1.0f32, f32::NAN, 2.0];
let err = extract_top_k_from_row_logits(&logits, 2).unwrap_err();
assert!(err.to_string().contains("not finite"), "got: {err}");
}
#[test]
fn adr_037_e4b10a_top_k_rejects_inf_logit_2026_05_22() {
let logits = vec![1.0f32, f32::INFINITY, 2.0];
let err = extract_top_k_from_row_logits(&logits, 2).unwrap_err();
assert!(err.to_string().contains("not finite"), "got: {err}");
}
#[test]
fn adr_037_e4b10a_top_k_deterministic_tie_break_by_smaller_token_2026_05_22() {
let logits = vec![5.0f32; 10];
let out = extract_top_k_from_row_logits(&logits, 2).unwrap();
assert_eq!(out.len(), 2);
let mut tokens: Vec<u32> = out.iter().map(|c| c.token).collect();
tokens.sort();
assert_eq!(tokens, vec![0, 1]);
let expected = -(10.0f32).ln();
for c in &out {
assert!((c.log_prob - expected).abs() < 1e-5);
}
}
#[test]
fn adr_037_e4b10a_top_k_passes_validate_candidates_2026_05_22() {
let logits: Vec<f32> = (0..1000).map(|i| (i as f32) * 0.01 - 5.0).collect();
let out = extract_top_k_from_row_logits(&logits, 10).unwrap();
validate_candidates(&out, 10).expect("Phase E4a contract");
}
#[test]
fn adr_037_e4b10a_gate_top_k_preserves_order_under_clamp_2026_05_22() {
let logits = vec![0.0f32, -3e38, -2e38];
let out = extract_top_k_from_row_logits(&logits, 2).unwrap();
assert_eq!(out.len(), 2);
assert_eq!(out[0].token, 0, "highest logit");
assert_eq!(
out[1].token, 2,
"next-highest is -2e38 (token 2), NOT -3e38 (token 1)"
);
}
#[test]
fn adr_037_e4b10a_top_k_clamps_extreme_log_probs_at_floor_2026_05_22() {
let logits = vec![100.0f32, 0.0, 0.0, 0.0];
let out = extract_top_k_from_row_logits(&logits, 4).unwrap();
for c in &out {
assert!(c.log_prob.is_finite(), "log_prob must be finite");
assert!(
c.log_prob >= LOG_PROB_FLOOR,
"log_prob {} below floor {}",
c.log_prob,
LOG_PROB_FLOOR
);
}
}
#[test]
fn adr_037_e4a_log_prob_floor_is_finite_2026_05_22() {
assert!(LOG_PROB_FLOOR.is_finite());
assert!(LOG_PROB_FLOOR < -1e30);
}
}