use amari_holographic::{BindingAlgebra, Resonator, ResonatorConfig};
use crate::error::{MinuetError, MinuetResult};
use crate::traits::{CleanupResult, RetrievalContext, Retriever};
pub struct ResonatorRetriever<A: BindingAlgebra> {
resonator: Option<Resonator<A>>,
config: ResonatorConfig,
}
impl<A: BindingAlgebra> Default for ResonatorRetriever<A> {
fn default() -> Self {
Self::new()
}
}
impl<A: BindingAlgebra> ResonatorRetriever<A> {
#[must_use]
pub fn new() -> Self {
Self {
resonator: None,
config: ResonatorConfig::default(),
}
}
#[must_use]
pub fn with_resonator(resonator: Resonator<A>) -> Self {
Self {
resonator: Some(resonator),
config: ResonatorConfig::default(),
}
}
pub fn from_symbols(symbols: Vec<A>) -> MinuetResult<Self> {
if symbols.is_empty() {
return Err(MinuetError::config("Codebook cannot be empty"));
}
let config = ResonatorConfig::default();
let resonator = Resonator::new(symbols, config.clone()).map_err(MinuetError::algebra)?;
Ok(Self {
resonator: Some(resonator),
config,
})
}
#[must_use]
pub fn with_config(mut self, config: ResonatorConfig) -> Self {
self.config = config;
self
}
#[must_use]
pub fn initial_temperature(mut self, temp: f64) -> Self {
self.config.initial_beta = temp;
self
}
#[must_use]
pub fn final_temperature(mut self, temp: f64) -> Self {
self.config.final_beta = temp;
self
}
#[must_use]
pub fn max_iterations(mut self, iters: usize) -> Self {
self.config.max_iterations = iters;
self
}
#[must_use]
pub fn with_temperature(mut self, temperature: &super::Temperature) -> Self {
self.config = temperature.to_resonator_config();
self.resonator = None;
self
}
#[must_use]
pub fn with_temperature_schedule(mut self, schedule: &super::TemperatureSchedule) -> Self {
if let (Some(&first), Some(&last)) = (
schedule.temperatures().first(),
schedule.temperatures().last(),
) {
self.config.initial_beta = first;
self.config.final_beta = last;
self.config.max_iterations = schedule.len();
self.resonator = None;
}
self
}
}
impl<A: BindingAlgebra> Retriever for ResonatorRetriever<A> {
type Algebra = A;
fn cleanup(&self, raw: &A, context: &RetrievalContext<A>) -> MinuetResult<CleanupResult<A>> {
let result = if let Some(ref resonator) = self.resonator {
resonator.cleanup(raw)
} else if let Some(ref codebook) = context.codebook {
if codebook.is_empty() {
return Ok(CleanupResult {
value: raw.clone(),
confidence: 0.5,
iterations: 0,
converged: false,
codebook_match: None,
});
}
let temp_resonator = Resonator::new(codebook.clone(), self.config.clone())
.map_err(MinuetError::algebra)?;
temp_resonator.cleanup(raw)
} else {
return Ok(CleanupResult {
value: raw.clone(),
confidence: 0.5,
iterations: 0,
converged: false,
codebook_match: None,
});
};
Ok(CleanupResult {
value: result.cleaned,
confidence: result.final_similarity,
iterations: result.iterations,
converged: result.converged,
codebook_match: Some(result.best_match_index),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::retrieval::{Temperature, TemperatureSchedule};
use amari_holographic::ProductCliffordAlgebra;
type TestAlgebra = ProductCliffordAlgebra<8>;
#[test]
fn cleanup_with_codebook() {
let symbols: Vec<TestAlgebra> = (0..5).map(|_| TestAlgebra::random_versor(2)).collect();
let retriever = ResonatorRetriever::from_symbols(symbols.clone()).unwrap();
let context = RetrievalContext::default();
let result = retriever.cleanup(&symbols[2], &context).unwrap();
assert!(result.converged);
assert!(result.confidence > 0.9);
assert!(result.codebook_match.is_some());
}
#[test]
fn cleanup_without_codebook_returns_raw() {
let retriever = ResonatorRetriever::<TestAlgebra>::new();
let raw = TestAlgebra::random_versor(2);
let context = RetrievalContext::default();
let result = retriever.cleanup(&raw, &context).unwrap();
assert!(result.value.similarity(&raw) > 0.99);
}
#[test]
fn cleanup_with_annealed_temperature() {
let symbols: Vec<TestAlgebra> = (0..5).map(|_| TestAlgebra::random_versor(2)).collect();
let retriever = ResonatorRetriever::new()
.with_temperature(&Temperature::annealed(1.0, 100.0, 50).unwrap());
let context = RetrievalContext::default().with_codebook(symbols.clone());
let result = retriever.cleanup(&symbols[2], &context).unwrap();
assert!(result.converged);
assert!(result.confidence > 0.9);
assert_eq!(result.codebook_match, Some(2));
}
#[test]
fn with_temperature_overwrites_config() {
let symbols: Vec<TestAlgebra> = (0..3).map(|_| TestAlgebra::random_versor(2)).collect();
let retriever = ResonatorRetriever::from_symbols(symbols).unwrap();
assert_eq!(retriever.config.max_iterations, 50);
let retriever = retriever.with_temperature(&Temperature::annealed(2.0, 50.0, 12).unwrap());
assert_eq!(retriever.config.initial_beta, 2.0);
assert_eq!(retriever.config.final_beta, 50.0);
assert_eq!(retriever.config.max_iterations, 12);
assert!(retriever.resonator.is_none());
}
#[test]
fn with_temperature_schedule_maps_endpoints() {
let retriever = ResonatorRetriever::<TestAlgebra>::new()
.with_temperature_schedule(&TemperatureSchedule::cosine(1.0, 10.0, 8));
assert!((retriever.config.initial_beta - 1.0).abs() < 1e-12);
assert!((retriever.config.final_beta - 10.0).abs() < 1e-12);
assert_eq!(retriever.config.max_iterations, 8);
}
#[test]
fn with_empty_schedule_is_noop() {
let retriever = ResonatorRetriever::<TestAlgebra>::new()
.with_temperature_schedule(&TemperatureSchedule::constant(1.0, 0));
assert_eq!(retriever.config.max_iterations, 50);
}
}