use crate::error::{InferenceError, InferenceResult};
use crate::sampling::{Sampler, SamplingConfig};
use kizzasi_model::AutoregressiveModel;
use scirs2_core::ndarray::Array1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EnsembleStrategy {
Average,
Weighted,
Voting,
ProductOfExperts,
}
#[derive(Debug, Clone)]
pub struct EnsembleConfig {
pub strategy: EnsembleStrategy,
pub weights: Option<Vec<f32>>,
pub normalize_outputs: bool,
pub temperature: f32,
}
impl Default for EnsembleConfig {
fn default() -> Self {
Self {
strategy: EnsembleStrategy::Average,
weights: None,
normalize_outputs: true,
temperature: 1.0,
}
}
}
impl EnsembleConfig {
pub fn new() -> Self {
Self::default()
}
pub fn strategy(mut self, strategy: EnsembleStrategy) -> Self {
self.strategy = strategy;
self
}
pub fn weights(mut self, weights: Vec<f32>) -> Self {
self.weights = Some(weights);
self
}
pub fn normalize_outputs(mut self, normalize: bool) -> Self {
self.normalize_outputs = normalize;
self
}
pub fn temperature(mut self, temp: f32) -> Self {
self.temperature = temp;
self
}
}
pub struct ModelEnsemble {
models: Vec<Box<dyn AutoregressiveModel>>,
config: EnsembleConfig,
sampler: Sampler,
}
impl ModelEnsemble {
pub fn new(
models: Vec<Box<dyn AutoregressiveModel>>,
config: EnsembleConfig,
) -> InferenceResult<Self> {
if models.is_empty() {
return Err(InferenceError::ForwardError(
"Ensemble must contain at least one model".to_string(),
));
}
if let Some(ref weights) = config.weights {
if weights.len() != models.len() {
return Err(InferenceError::DimensionMismatch {
expected: models.len(),
got: weights.len(),
});
}
let sum: f32 = weights.iter().sum();
if (sum - 1.0).abs() > 1e-6 {
return Err(InferenceError::ForwardError(format!(
"Ensemble weights must sum to 1.0, got {}",
sum
)));
}
}
let sampler_config = SamplingConfig::new().temperature(config.temperature);
let sampler = Sampler::new(sampler_config);
Ok(Self {
models,
config,
sampler,
})
}
pub fn num_models(&self) -> usize {
self.models.len()
}
pub fn step(&mut self, input: &Array1<f32>) -> InferenceResult<Array1<f32>> {
let mut predictions = Vec::with_capacity(self.models.len());
for model in &mut self.models {
let pred = model
.step(input)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
predictions.push(pred);
}
self.combine_predictions(&predictions)
}
fn combine_predictions(&mut self, predictions: &[Array1<f32>]) -> InferenceResult<Array1<f32>> {
if predictions.is_empty() {
return Err(InferenceError::ForwardError(
"No predictions to combine".to_string(),
));
}
let output_dim = predictions[0].len();
for pred in predictions {
if pred.len() != output_dim {
return Err(InferenceError::DimensionMismatch {
expected: output_dim,
got: pred.len(),
});
}
}
match self.config.strategy {
EnsembleStrategy::Average => self.combine_average(predictions, output_dim),
EnsembleStrategy::Weighted => self.combine_weighted(predictions, output_dim),
EnsembleStrategy::Voting => self.combine_voting(predictions),
EnsembleStrategy::ProductOfExperts => {
self.combine_product_of_experts(predictions, output_dim)
}
}
}
fn combine_average(
&self,
predictions: &[Array1<f32>],
output_dim: usize,
) -> InferenceResult<Array1<f32>> {
let mut combined = Array1::zeros(output_dim);
let n = predictions.len() as f32;
for pred in predictions {
combined += pred;
}
combined /= n;
if self.config.normalize_outputs {
combined = self.normalize(&combined);
}
Ok(combined)
}
fn combine_weighted(
&self,
predictions: &[Array1<f32>],
output_dim: usize,
) -> InferenceResult<Array1<f32>> {
let weights = self.config.weights.as_ref().ok_or_else(|| {
InferenceError::ForwardError("Weights not provided for weighted ensemble".to_string())
})?;
let mut combined = Array1::zeros(output_dim);
for (pred, &weight) in predictions.iter().zip(weights.iter()) {
combined += &(pred * weight);
}
if self.config.normalize_outputs {
combined = self.normalize(&combined);
}
Ok(combined)
}
fn combine_voting(&mut self, predictions: &[Array1<f32>]) -> InferenceResult<Array1<f32>> {
let votes: Vec<usize> = predictions
.iter()
.map(|pred| {
pred.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(idx, _)| idx)
.unwrap_or(0)
})
.collect();
let output_dim = predictions[0].len();
let mut vote_counts = vec![0usize; output_dim];
for &vote in &votes {
if vote < output_dim {
vote_counts[vote] += 1;
}
}
let total_votes = votes.len() as f32;
let combined = Array1::from_vec(
vote_counts
.iter()
.map(|&count| count as f32 / total_votes)
.collect(),
);
Ok(combined)
}
fn combine_product_of_experts(
&self,
predictions: &[Array1<f32>],
output_dim: usize,
) -> InferenceResult<Array1<f32>> {
let mut combined = Array1::ones(output_dim);
for pred in predictions {
let normalized = self.softmax(pred);
combined *= &normalized;
}
let sum: f32 = combined.sum();
if sum > 0.0 {
combined /= sum;
}
Ok(combined)
}
fn normalize(&self, output: &Array1<f32>) -> Array1<f32> {
self.softmax(output)
}
fn softmax(&self, x: &Array1<f32>) -> Array1<f32> {
let max_x = x.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let exp_x = x.mapv(|v| (v - max_x).exp());
let sum_exp: f32 = exp_x.sum();
if sum_exp > 0.0 {
exp_x / sum_exp
} else {
Array1::from_elem(x.len(), 1.0 / x.len() as f32)
}
}
pub fn config(&self) -> &EnsembleConfig {
&self.config
}
pub fn sampler_mut(&mut self) -> &mut Sampler {
&mut self.sampler
}
}
pub struct EnsembleBuilder {
models: Vec<Box<dyn AutoregressiveModel>>,
config: EnsembleConfig,
}
impl EnsembleBuilder {
pub fn new() -> Self {
Self {
models: Vec::new(),
config: EnsembleConfig::default(),
}
}
pub fn add_model(mut self, model: Box<dyn AutoregressiveModel>) -> Self {
self.models.push(model);
self
}
pub fn add_models(mut self, models: Vec<Box<dyn AutoregressiveModel>>) -> Self {
self.models.extend(models);
self
}
pub fn strategy(mut self, strategy: EnsembleStrategy) -> Self {
self.config.strategy = strategy;
self
}
pub fn weights(mut self, weights: Vec<f32>) -> Self {
self.config.weights = Some(weights);
self
}
pub fn temperature(mut self, temp: f32) -> Self {
self.config.temperature = temp;
self
}
pub fn build(self) -> InferenceResult<ModelEnsemble> {
ModelEnsemble::new(self.models, self.config)
}
}
impl Default for EnsembleBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use kizzasi_model::s4::{S4Config, S4D};
#[test]
fn test_ensemble_creation() {
let model1 = create_test_model();
let model2 = create_test_model();
let ensemble = EnsembleBuilder::new()
.add_model(Box::new(model1))
.add_model(Box::new(model2))
.build();
assert!(ensemble.is_ok());
let ensemble = ensemble.unwrap();
assert_eq!(ensemble.num_models(), 2);
}
#[test]
fn test_ensemble_average() {
let model1 = create_test_model();
let model2 = create_test_model();
let mut ensemble = EnsembleBuilder::new()
.add_model(Box::new(model1))
.add_model(Box::new(model2))
.strategy(EnsembleStrategy::Average)
.build()
.unwrap();
let input = Array1::from_vec(vec![0.5]);
let output = ensemble.step(&input);
assert!(output.is_ok());
}
#[test]
fn test_ensemble_weighted() {
let model1 = create_test_model();
let model2 = create_test_model();
let mut ensemble = EnsembleBuilder::new()
.add_model(Box::new(model1))
.add_model(Box::new(model2))
.strategy(EnsembleStrategy::Weighted)
.weights(vec![0.7, 0.3])
.build()
.unwrap();
let input = Array1::from_vec(vec![0.5]);
let output = ensemble.step(&input);
assert!(output.is_ok());
}
#[test]
fn test_ensemble_voting() {
let model1 = create_test_model();
let model2 = create_test_model();
let model3 = create_test_model();
let mut ensemble = EnsembleBuilder::new()
.add_model(Box::new(model1))
.add_model(Box::new(model2))
.add_model(Box::new(model3))
.strategy(EnsembleStrategy::Voting)
.build()
.unwrap();
let input = Array1::from_vec(vec![0.5]);
let output = ensemble.step(&input);
assert!(output.is_ok());
}
#[test]
fn test_invalid_weights() {
let model1 = create_test_model();
let model2 = create_test_model();
let result = EnsembleBuilder::new()
.add_model(Box::new(model1))
.add_model(Box::new(model2))
.strategy(EnsembleStrategy::Weighted)
.weights(vec![0.5, 0.6]) .build();
assert!(result.is_err());
}
fn create_test_model() -> S4D {
let config = S4Config::new()
.input_dim(1)
.hidden_dim(64)
.state_dim(16)
.num_layers(2)
.diagonal(true);
S4D::new(config).unwrap()
}
}