Skip to main content

lean_ctx/core/wasserstein/
allocator.rs

1//! Wasserstein/Sinkhorn token budget allocation across context files.
2//!
3//! Distributes a fixed token budget proportionally to relevance scores using
4//! optimal transport from a single supply node to per-file demand nodes.
5
6use super::transport::sinkhorn_plan;
7
8/// Allocation result for a single file or context chunk.
9#[derive(Debug, Clone)]
10pub struct TokenAllocation {
11    /// File path or chunk identifier.
12    pub target: String,
13    /// Allocated token budget.
14    pub tokens: usize,
15    /// Fraction of the total budget, in the range `0.0..=1.0`.
16    pub fraction: f64,
17}
18
19/// Allocates a token budget across files according to their relevance scores.
20///
21/// Each input tuple is `(path, current_tokens, relevance_score)`. Relevance is
22/// clamped to `0.0..=1.0`; when every score is zero, the budget is shared
23/// uniformly. `current_tokens` is retained for callers that use it to describe
24/// source size, but it does not cap the requested context budget.
25pub(crate) fn allocate_budget(
26    files: &[(&str, usize, f64)],
27    total_budget: usize,
28) -> Vec<TokenAllocation> {
29    if files.is_empty() {
30        return Vec::new();
31    }
32
33    let relevance: Vec<f64> = files
34        .iter()
35        .map(|(_, _, score)| sanitize_relevance(*score))
36        .collect();
37    let relevance_sum: f64 = relevance.iter().sum();
38    let weights = if relevance_sum > 0.0 {
39        relevance.clone()
40    } else {
41        vec![1.0; files.len()]
42    };
43    let weight_sum: f64 = weights.iter().sum();
44    let demand: Vec<f64> = weights
45        .iter()
46        .map(|weight| total_budget as f64 * weight / weight_sum)
47        .collect();
48    let cost_matrix = vec![relevance.iter().map(|score| 1.0 - score).collect()];
49    let plan = sinkhorn_plan(&[total_budget as f64], &demand, &cost_matrix, 0.1, 50);
50
51    let mut transported = vec![0.0; files.len()];
52    for (_, target_idx, amount) in &plan {
53        if let Some(target) = transported.get_mut(*target_idx) {
54            *target += amount;
55        }
56    }
57    if plan.is_empty() || transported.iter().all(|amount| *amount == 0.0) {
58        transported = demand;
59    }
60    let integer_allocations = round_allocations(&transported, total_budget);
61
62    files
63        .iter()
64        .zip(integer_allocations)
65        .map(|((target, _, _), tokens)| TokenAllocation {
66            target: (*target).to_owned(),
67            tokens,
68            fraction: if total_budget == 0 {
69                0.0
70            } else {
71                tokens as f64 / total_budget as f64
72            },
73        })
74        .collect()
75}
76
77fn sanitize_relevance(score: f64) -> f64 {
78    if score.is_finite() {
79        score.clamp(0.0, 1.0)
80    } else {
81        0.0
82    }
83}
84
85fn round_allocations(values: &[f64], budget: usize) -> Vec<usize> {
86    let mut rounded: Vec<usize> = values
87        .iter()
88        .map(|value| value.max(0.0).floor() as usize)
89        .collect();
90    let assigned = rounded.iter().sum::<usize>().min(budget);
91    let remaining = budget - assigned;
92    let mut by_remainder: Vec<usize> = (0..values.len()).collect();
93    by_remainder.sort_by(|left, right| {
94        let left_remainder = values[*left] - values[*left].floor();
95        let right_remainder = values[*right] - values[*right].floor();
96        right_remainder
97            .total_cmp(&left_remainder)
98            .then_with(|| left.cmp(right))
99    });
100
101    for index in by_remainder.into_iter().take(remaining) {
102        rounded[index] = rounded[index].saturating_add(1);
103    }
104    rounded
105}
106
107#[cfg(test)]
108mod tests {
109    use super::allocate_budget;
110
111    #[test]
112    fn single_file_gets_full_budget() {
113        let allocations = allocate_budget(&[("src/lib.rs", 20, 0.2)], 100);
114        assert_eq!(allocations[0].tokens, 100);
115        assert_eq!(allocations[0].fraction, 1.0);
116    }
117
118    #[test]
119    fn irrelevant_file_gets_minimum() {
120        let allocations = allocate_budget(
121            &[("relevant.rs", 100, 1.0), ("irrelevant.rs", 100, 0.0)],
122            100,
123        );
124        assert_eq!(allocations[1].tokens, 0);
125    }
126
127    #[test]
128    fn allocation_sums_to_budget() {
129        let allocations = allocate_budget(
130            &[("a.rs", 100, 0.7), ("b.rs", 80, 0.2), ("c.rs", 40, 0.1)],
131            101,
132        );
133        assert_eq!(
134            allocations.iter().map(|entry| entry.tokens).sum::<usize>(),
135            101
136        );
137    }
138
139    #[test]
140    fn higher_relevance_gets_more_tokens() {
141        let allocations = allocate_budget(&[("high.rs", 100, 0.9), ("low.rs", 100, 0.1)], 100);
142        assert!(allocations[0].tokens > allocations[1].tokens);
143    }
144}