Skip to main content

rag_plusplus_core/trajectory/
salience.rs

1//! Salience Scoring System
2//!
3//! Computes importance weights for episodes, enabling bounded forgetting
4//! and prioritized retrieval.
5//!
6//! Per SPECIFICATION_HARD.md Section 9:
7//!
8//! ```text
9//! salience_score = weighted_sum([
10//!     user_feedback,           // thumbs_up = +0.35, thumbs_down = -0.35
11//!     phase_transition,        // +0.10 at transitions
12//!     downstream_references,   // +0.05 per ref, capped at +0.15
13//!     novelty,                 // +0.10 max based on embedding distance
14//! ])
15//! ```
16//!
17//! Scores are bounded to [0, 1].
18//!
19//! # Debug Logging
20//!
21//! Enable the `salience-debug` feature to log salience distribution statistics
22//! for tuning and analysis.
23
24use crate::distance::cosine_similarity;
25
26/// User feedback type.
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum Feedback {
29    ThumbsUp,
30    ThumbsDown,
31    None,
32}
33
34/// Factors contributing to salience score.
35#[derive(Debug, Clone, Default)]
36pub struct SalienceFactors {
37    /// Episode identifier
38    pub turn_id: u64,
39    /// User feedback on this episode
40    pub feedback: Option<Feedback>,
41    /// Whether this is a phase transition point
42    pub is_phase_transition: bool,
43    /// Number of times referenced by later episodes
44    pub reference_count: usize,
45    /// Embedding of this episode (for novelty)
46    pub embedding: Option<Vec<f32>>,
47    /// Episode role
48    pub role: String,
49    /// Content length
50    pub content_length: usize,
51}
52
53/// Result of salience computation for an episode.
54#[derive(Debug, Clone)]
55pub struct TurnSalience {
56    pub turn_id: u64,
57    pub score: f32,
58    pub feedback_contribution: f32,
59    pub phase_contribution: f32,
60    pub reference_contribution: f32,
61    pub novelty_contribution: f32,
62}
63
64/// Corpus-level salience statistics.
65#[derive(Debug, Clone, Default)]
66pub struct CorpusSalienceStats {
67    pub total_turns: usize,
68    pub mean_salience: f32,
69    pub std_salience: f32,
70    pub min_salience: f32,
71    pub max_salience: f32,
72    pub high_salience_count: usize,  // > 0.7
73    pub low_salience_count: usize,   // < 0.3
74}
75
76/// Configuration for salience scoring.
77#[derive(Debug, Clone)]
78pub struct SalienceConfig {
79    /// Base salience score (neutral)
80    pub base_score: f32,
81    /// Boost for thumbs_up
82    pub feedback_up_boost: f32,
83    /// Penalty for thumbs_down
84    pub feedback_down_penalty: f32,
85    /// Boost for phase transitions
86    pub phase_transition_boost: f32,
87    /// Boost per reference (capped)
88    pub reference_boost_per: f32,
89    /// Maximum reference contribution
90    pub reference_boost_cap: f32,
91    /// Maximum novelty contribution
92    pub novelty_boost_max: f32,
93    /// Size of recent window for novelty computation
94    pub novelty_window_size: usize,
95    /// Version string
96    pub version: String,
97}
98
99impl Default for SalienceConfig {
100    fn default() -> Self {
101        // Per SPECIFICATION_HARD.md Section 9
102        Self {
103            base_score: 0.5,
104            feedback_up_boost: 0.35,
105            feedback_down_penalty: 0.35,
106            phase_transition_boost: 0.10,
107            reference_boost_per: 0.05,
108            reference_boost_cap: 0.15,
109            novelty_boost_max: 0.10,
110            novelty_window_size: 10,
111            version: "v1.0".to_string(),
112        }
113    }
114}
115
116/// Salience scoring engine.
117#[derive(Debug, Clone)]
118pub struct SalienceScorer {
119    config: SalienceConfig,
120}
121
122impl SalienceScorer {
123    /// Create with default configuration.
124    pub fn new() -> Self {
125        Self {
126            config: SalienceConfig::default(),
127        }
128    }
129
130    /// Create with custom configuration.
131    pub fn with_config(config: SalienceConfig) -> Self {
132        Self { config }
133    }
134
135    /// Compute salience for a single episode.
136    ///
137    /// # Arguments
138    ///
139    /// * `factors` - Features of the episode
140    /// * `recent_embeddings` - Embeddings of recent episodes (for novelty)
141    ///
142    /// # Returns
143    ///
144    /// TurnSalience with score and contribution breakdown.
145    pub fn score_single(
146        &self,
147        factors: &SalienceFactors,
148        recent_embeddings: Option<&[Vec<f32>]>,
149    ) -> TurnSalience {
150        let mut score = self.config.base_score;
151
152        // Feedback contribution
153        let feedback_contribution = match factors.feedback {
154            Some(Feedback::ThumbsUp) => self.config.feedback_up_boost,
155            Some(Feedback::ThumbsDown) => -self.config.feedback_down_penalty,
156            Some(Feedback::None) | None => 0.0,
157        };
158        score += feedback_contribution;
159
160        // Phase transition contribution
161        let phase_contribution = if factors.is_phase_transition {
162            self.config.phase_transition_boost
163        } else {
164            0.0
165        };
166        score += phase_contribution;
167
168        // Reference contribution (capped)
169        let reference_contribution = (factors.reference_count as f32 * self.config.reference_boost_per)
170            .min(self.config.reference_boost_cap);
171        score += reference_contribution;
172
173        // Novelty contribution
174        let novelty_contribution = self.compute_novelty(factors, recent_embeddings);
175        score += novelty_contribution;
176
177        // Clamp to [0, 1]
178        score = score.clamp(0.0, 1.0);
179
180        TurnSalience {
181            turn_id: factors.turn_id,
182            score,
183            feedback_contribution,
184            phase_contribution,
185            reference_contribution,
186            novelty_contribution,
187        }
188    }
189
190    /// Compute novelty based on embedding distance from recent window.
191    fn compute_novelty(
192        &self,
193        factors: &SalienceFactors,
194        recent_embeddings: Option<&[Vec<f32>]>,
195    ) -> f32 {
196        let embedding = match &factors.embedding {
197            Some(e) => e,
198            None => return 0.0,
199        };
200
201        let recent = match recent_embeddings {
202            Some(r) if !r.is_empty() => r,
203            _ => return self.config.novelty_boost_max, // First episode is maximally novel
204        };
205
206        // Compute average similarity to recent episodes
207        let similarities: Vec<f32> = recent
208            .iter()
209            .take(self.config.novelty_window_size)
210            .map(|other| cosine_similarity(embedding, other))
211            .collect();
212
213        if similarities.is_empty() {
214            return self.config.novelty_boost_max;
215        }
216
217        let avg_similarity = similarities.iter().sum::<f32>() / similarities.len() as f32;
218
219        // Novelty = 1 - similarity, scaled to max boost
220        let novelty = 1.0 - avg_similarity;
221        (novelty * self.config.novelty_boost_max).clamp(0.0, self.config.novelty_boost_max)
222    }
223
224    /// Compute salience for all episodes in a corpus.
225    ///
226    /// This is the high-performance batch version that processes
227    /// the entire corpus efficiently.
228    ///
229    /// # Arguments
230    ///
231    /// * `turns` - All episodes with their factors
232    ///
233    /// # Returns
234    ///
235    /// Vector of salience scores for each episode.
236    pub fn score_corpus(&self, turns: &[SalienceFactors]) -> Vec<TurnSalience> {
237        let mut results = Vec::with_capacity(turns.len());
238        let mut recent_embeddings: Vec<&Vec<f32>> = Vec::with_capacity(self.config.novelty_window_size);
239
240        for factors in turns {
241            // Build recent embeddings window
242            let recent: Vec<Vec<f32>> = recent_embeddings
243                .iter()
244                .map(|e| (*e).clone())
245                .collect();
246
247            let salience = self.score_single(
248                factors,
249                if recent.is_empty() { None } else { Some(&recent) },
250            );
251
252            results.push(salience);
253
254            // Update recent window
255            if let Some(ref emb) = factors.embedding {
256                recent_embeddings.push(emb);
257                if recent_embeddings.len() > self.config.novelty_window_size {
258                    recent_embeddings.remove(0);
259                }
260            }
261        }
262
263        // Debug logging for salience distribution tuning
264        #[cfg(feature = "salience-debug")]
265        self.log_salience_distribution(&results);
266
267        results
268    }
269
270    /// Log salience distribution statistics for tuning.
271    ///
272    /// Only available with `salience-debug` feature enabled.
273    #[cfg(feature = "salience-debug")]
274    fn log_salience_distribution(&self, saliences: &[TurnSalience]) {
275        if saliences.is_empty() {
276            return;
277        }
278
279        let scores: Vec<f32> = saliences.iter().map(|s| s.score).collect();
280        let n = scores.len() as f32;
281        let mean = scores.iter().sum::<f32>() / n;
282        let variance = scores.iter().map(|s| (s - mean).powi(2)).sum::<f32>() / n;
283        let std = variance.sqrt();
284        let min = scores.iter().cloned().fold(f32::INFINITY, f32::min);
285        let max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
286
287        // Compute percentiles
288        let mut sorted = scores.clone();
289        sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
290        let p25 = sorted.get((n * 0.25) as usize).copied().unwrap_or(0.0);
291        let p50 = sorted.get((n * 0.50) as usize).copied().unwrap_or(0.0);
292        let p75 = sorted.get((n * 0.75) as usize).copied().unwrap_or(0.0);
293
294        // Count by range
295        let low = scores.iter().filter(|&&s| s < 0.3).count();
296        let mid = scores.iter().filter(|&&s| s >= 0.3 && s <= 0.7).count();
297        let high = scores.iter().filter(|&&s| s > 0.7).count();
298
299        // Log with tracing
300        tracing::debug!(
301            n = scores.len(),
302            mean = %format!("{:.3}", mean),
303            std = %format!("{:.3}", std),
304            min = %format!("{:.3}", min),
305            max = %format!("{:.3}", max),
306            p25 = %format!("{:.3}", p25),
307            p50 = %format!("{:.3}", p50),
308            p75 = %format!("{:.3}", p75),
309            low_count = low,
310            mid_count = mid,
311            high_count = high,
312            "Salience distribution"
313        );
314    }
315
316    /// Compute corpus-level statistics.
317    pub fn compute_stats(&self, saliences: &[TurnSalience]) -> CorpusSalienceStats {
318        if saliences.is_empty() {
319            return CorpusSalienceStats::default();
320        }
321
322        let scores: Vec<f32> = saliences.iter().map(|s| s.score).collect();
323        let n = scores.len();
324
325        let mean = scores.iter().sum::<f32>() / n as f32;
326        let variance = scores.iter()
327            .map(|s| (s - mean).powi(2))
328            .sum::<f32>() / n as f32;
329        let std = variance.sqrt();
330
331        let min = scores.iter().cloned().fold(f32::INFINITY, f32::min);
332        let max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
333
334        let high_count = scores.iter().filter(|&&s| s > 0.7).count();
335        let low_count = scores.iter().filter(|&&s| s < 0.3).count();
336
337        CorpusSalienceStats {
338            total_turns: n,
339            mean_salience: mean,
340            std_salience: std,
341            min_salience: min,
342            max_salience: max,
343            high_salience_count: high_count,
344            low_salience_count: low_count,
345        }
346    }
347
348    /// Normalize scores to have mean 0.5 and bounded to [0, 1].
349    ///
350    /// Useful when raw scores cluster too tightly.
351    pub fn normalize_scores(&self, saliences: &mut [TurnSalience]) {
352        if saliences.is_empty() {
353            return;
354        }
355
356        let scores: Vec<f32> = saliences.iter().map(|s| s.score).collect();
357        let mean = scores.iter().sum::<f32>() / scores.len() as f32;
358        let std = {
359            let variance = scores.iter()
360                .map(|s| (s - mean).powi(2))
361                .sum::<f32>() / scores.len() as f32;
362            variance.sqrt().max(0.01) // Avoid division by zero
363        };
364
365        // Z-score normalization centered at 0.5
366        for s in saliences.iter_mut() {
367            let z = (s.score - mean) / std;
368            s.score = (0.5 + z * 0.2).clamp(0.0, 1.0);
369        }
370    }
371
372    /// Get version string.
373    #[inline]
374    pub fn version(&self) -> &str {
375        &self.config.version
376    }
377}
378
379impl Default for SalienceScorer {
380    fn default() -> Self {
381        Self::new()
382    }
383}
384
385#[cfg(test)]
386mod tests {
387    use super::*;
388
389    fn make_factors(turn_id: u64, feedback: Option<Feedback>) -> SalienceFactors {
390        SalienceFactors {
391            turn_id,
392            feedback,
393            is_phase_transition: false,
394            reference_count: 0,
395            embedding: None,
396            role: "assistant".to_string(),
397            content_length: 100,
398        }
399    }
400
401    #[test]
402    fn test_base_score() {
403        let scorer = SalienceScorer::new();
404        let factors = make_factors(1, None);
405        let result = scorer.score_single(&factors, None);
406
407        assert!((result.score - 0.5).abs() < 0.01);
408    }
409
410    #[test]
411    fn test_thumbs_up_boost() {
412        let scorer = SalienceScorer::new();
413        let factors = make_factors(1, Some(Feedback::ThumbsUp));
414        let result = scorer.score_single(&factors, None);
415
416        assert!((result.score - 0.85).abs() < 0.01); // 0.5 + 0.35
417        assert!((result.feedback_contribution - 0.35).abs() < 0.01);
418    }
419
420    #[test]
421    fn test_thumbs_down_penalty() {
422        let scorer = SalienceScorer::new();
423        let factors = make_factors(1, Some(Feedback::ThumbsDown));
424        let result = scorer.score_single(&factors, None);
425
426        assert!((result.score - 0.15).abs() < 0.01); // 0.5 - 0.35
427    }
428
429    #[test]
430    fn test_phase_transition_boost() {
431        let scorer = SalienceScorer::new();
432        let mut factors = make_factors(1, None);
433        factors.is_phase_transition = true;
434        let result = scorer.score_single(&factors, None);
435
436        assert!((result.score - 0.6).abs() < 0.01); // 0.5 + 0.10
437        assert!((result.phase_contribution - 0.10).abs() < 0.01);
438    }
439
440    #[test]
441    fn test_reference_boost_capped() {
442        let scorer = SalienceScorer::new();
443        let mut factors = make_factors(1, None);
444        factors.reference_count = 10; // Would be 0.5 without cap
445        let result = scorer.score_single(&factors, None);
446
447        assert!((result.reference_contribution - 0.15).abs() < 0.01); // Capped
448        assert!((result.score - 0.65).abs() < 0.01); // 0.5 + 0.15
449    }
450
451    #[test]
452    fn test_score_clamping() {
453        let scorer = SalienceScorer::new();
454        let mut factors = make_factors(1, Some(Feedback::ThumbsUp));
455        factors.is_phase_transition = true;
456        factors.reference_count = 5;
457        // Would be 0.5 + 0.35 + 0.10 + 0.15 + novelty = > 1.0
458        let result = scorer.score_single(&factors, None);
459
460        assert!(result.score <= 1.0);
461        assert!(result.score >= 0.0);
462    }
463
464    #[test]
465    fn test_corpus_scoring() {
466        let scorer = SalienceScorer::new();
467        let turns = vec![
468            make_factors(1, None),
469            make_factors(2, Some(Feedback::ThumbsUp)),
470            make_factors(3, Some(Feedback::ThumbsDown)),
471        ];
472
473        let results = scorer.score_corpus(&turns);
474        assert_eq!(results.len(), 3);
475
476        // Verify ordering
477        assert!(results[1].score > results[0].score); // ThumbsUp > None
478        assert!(results[0].score > results[2].score); // None > ThumbsDown
479    }
480
481    #[test]
482    fn test_corpus_stats() {
483        let scorer = SalienceScorer::new();
484        let saliences = vec![
485            TurnSalience { turn_id: 1, score: 0.2, feedback_contribution: 0.0, phase_contribution: 0.0, reference_contribution: 0.0, novelty_contribution: 0.0 },
486            TurnSalience { turn_id: 2, score: 0.5, feedback_contribution: 0.0, phase_contribution: 0.0, reference_contribution: 0.0, novelty_contribution: 0.0 },
487            TurnSalience { turn_id: 3, score: 0.8, feedback_contribution: 0.0, phase_contribution: 0.0, reference_contribution: 0.0, novelty_contribution: 0.0 },
488        ];
489
490        let stats = scorer.compute_stats(&saliences);
491        assert_eq!(stats.total_turns, 3);
492        assert!((stats.mean_salience - 0.5).abs() < 0.01);
493        assert!(stats.min_salience < 0.3);
494        assert!(stats.max_salience > 0.7);
495    }
496}