use serde::{Deserialize, Serialize};
use super::aggregation::FitnessAggregator;
use super::evaluator::{Candidate, CandidateId, EvaluationRequest, EvaluationResponse};
use crate::genome::traits::EvolutionaryGenome;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum EvaluationMode {
Rating,
Pairwise,
BatchSelection,
Adaptive,
}
impl EvaluationMode {
pub fn description(&self) -> &'static str {
match self {
Self::Rating => "Rate each candidate on a numeric scale",
Self::Pairwise => "Compare pairs and select the better one",
Self::BatchSelection => "Select favorites from a batch",
Self::Adaptive => "System adapts evaluation method automatically",
}
}
}
impl Default for EvaluationMode {
fn default() -> Self {
Self::Rating
}
}
pub trait InteractiveFitness: Send + Sync {
type Genome: EvolutionaryGenome;
fn evaluation_mode(&self) -> EvaluationMode;
fn request_evaluation(
&self,
candidates: &[Candidate<Self::Genome>],
) -> EvaluationRequest<Self::Genome>;
fn process_response(
&mut self,
response: EvaluationResponse,
aggregator: &mut FitnessAggregator,
) -> Vec<(CandidateId, f64)>;
fn on_generation_start(&mut self, _generation: usize, _population_size: usize) {}
fn on_evaluation_skipped(&mut self) {}
}
#[derive(Clone, Debug)]
pub struct DefaultInteractiveFitness<G>
where
G: EvolutionaryGenome,
{
mode: EvaluationMode,
batch_size: usize,
select_count: usize,
_marker: std::marker::PhantomData<G>,
}
impl<G> DefaultInteractiveFitness<G>
where
G: EvolutionaryGenome,
{
pub fn new(mode: EvaluationMode) -> Self {
Self {
mode,
batch_size: 6,
select_count: 2,
_marker: std::marker::PhantomData,
}
}
pub fn with_batch_size(mut self, size: usize) -> Self {
self.batch_size = size;
self
}
pub fn with_select_count(mut self, count: usize) -> Self {
self.select_count = count;
self
}
}
impl<G> Default for DefaultInteractiveFitness<G>
where
G: EvolutionaryGenome,
{
fn default() -> Self {
Self::new(EvaluationMode::Rating)
}
}
impl<G> InteractiveFitness for DefaultInteractiveFitness<G>
where
G: EvolutionaryGenome + Clone + Send + Sync,
{
type Genome = G;
fn evaluation_mode(&self) -> EvaluationMode {
self.mode
}
fn request_evaluation(
&self,
candidates: &[Candidate<Self::Genome>],
) -> EvaluationRequest<Self::Genome> {
match self.mode {
EvaluationMode::Rating => EvaluationRequest::rate(candidates.to_vec()),
EvaluationMode::Pairwise => {
if candidates.len() >= 2 {
EvaluationRequest::compare(candidates[0].clone(), candidates[1].clone())
} else if candidates.len() == 1 {
EvaluationRequest::rate(candidates.to_vec())
} else {
EvaluationRequest::rate(vec![])
}
}
EvaluationMode::BatchSelection => {
let batch: Vec<_> = candidates.iter().take(self.batch_size).cloned().collect();
EvaluationRequest::select_from_batch(batch, self.select_count)
}
EvaluationMode::Adaptive => {
EvaluationRequest::rate(candidates.to_vec())
}
}
}
fn process_response(
&mut self,
response: EvaluationResponse,
aggregator: &mut FitnessAggregator,
) -> Vec<(CandidateId, f64)> {
aggregator.process_response(&response)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::genome::real_vector::RealVector;
use crate::interactive::aggregation::AggregationModel;
#[test]
fn test_evaluation_mode_default() {
assert_eq!(EvaluationMode::default(), EvaluationMode::Rating);
}
#[test]
fn test_evaluation_mode_description() {
assert!(!EvaluationMode::Rating.description().is_empty());
assert!(!EvaluationMode::Pairwise.description().is_empty());
assert!(!EvaluationMode::BatchSelection.description().is_empty());
assert!(!EvaluationMode::Adaptive.description().is_empty());
}
#[test]
fn test_default_interactive_fitness_rating() {
let fitness: DefaultInteractiveFitness<RealVector> =
DefaultInteractiveFitness::new(EvaluationMode::Rating);
let c1 = Candidate::new(CandidateId(0), RealVector::new(vec![1.0]));
let c2 = Candidate::new(CandidateId(1), RealVector::new(vec![2.0]));
let request = fitness.request_evaluation(&[c1, c2]);
match request {
EvaluationRequest::RateCandidates { candidates, .. } => {
assert_eq!(candidates.len(), 2);
}
_ => panic!("Expected RateCandidates request"),
}
}
#[test]
fn test_default_interactive_fitness_pairwise() {
let fitness: DefaultInteractiveFitness<RealVector> =
DefaultInteractiveFitness::new(EvaluationMode::Pairwise);
let c1 = Candidate::new(CandidateId(0), RealVector::new(vec![1.0]));
let c2 = Candidate::new(CandidateId(1), RealVector::new(vec![2.0]));
let request = fitness.request_evaluation(&[c1, c2]);
match request {
EvaluationRequest::PairwiseComparison { .. } => {}
_ => panic!("Expected PairwiseComparison request"),
}
}
#[test]
fn test_default_interactive_fitness_batch() {
let fitness: DefaultInteractiveFitness<RealVector> =
DefaultInteractiveFitness::new(EvaluationMode::BatchSelection)
.with_batch_size(4)
.with_select_count(2);
let candidates: Vec<_> = (0..6)
.map(|i| Candidate::new(CandidateId(i), RealVector::new(vec![i as f64])))
.collect();
let request = fitness.request_evaluation(&candidates);
match request {
EvaluationRequest::BatchSelection {
candidates,
select_count,
..
} => {
assert_eq!(candidates.len(), 4); assert_eq!(select_count, 2);
}
_ => panic!("Expected BatchSelection request"),
}
}
#[test]
fn test_default_interactive_fitness_process_response() {
let mut fitness: DefaultInteractiveFitness<RealVector> =
DefaultInteractiveFitness::new(EvaluationMode::Rating);
let mut aggregator = FitnessAggregator::new(AggregationModel::DirectRating {
default_rating: 5.0,
});
let response =
EvaluationResponse::ratings(vec![(CandidateId(0), 8.0), (CandidateId(1), 6.0)]);
let updated = fitness.process_response(response, &mut aggregator);
assert_eq!(updated.len(), 2);
}
}