lean_ctx/core/wasserstein/
allocator.rs1use super::transport::sinkhorn_plan;
7
8#[derive(Debug, Clone)]
10pub struct TokenAllocation {
11 pub target: String,
13 pub tokens: usize,
15 pub fraction: f64,
17}
18
19pub(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}