use crate::error::{KizzasiError, KizzasiResult};
use crate::plugin::{Plugin, PluginManager};
use kizzasi_core::{KizzasiConfig, ModelType, SelectiveSSM, SignalPredictor, StateSpaceModel};
use scirs2_core::ndarray::{Array1, Array2};
#[cfg(feature = "logic")]
use kizzasi_logic::{ConstrainedInference, GuardrailSet};
impl From<f32> for SignalInput {
fn from(value: f32) -> Self {
SignalInput(Array1::from_vec(vec![value]))
}
}
impl From<Vec<f32>> for SignalInput {
fn from(value: Vec<f32>) -> Self {
SignalInput(Array1::from_vec(value))
}
}
impl From<&[f32]> for SignalInput {
fn from(value: &[f32]) -> Self {
SignalInput(Array1::from_vec(value.to_vec()))
}
}
impl<const N: usize> From<[f32; N]> for SignalInput {
fn from(value: [f32; N]) -> Self {
SignalInput(Array1::from_vec(value.to_vec()))
}
}
impl From<Array1<f32>> for SignalInput {
fn from(value: Array1<f32>) -> Self {
SignalInput(value)
}
}
#[derive(Debug, Clone)]
pub struct SignalInput(pub Array1<f32>);
impl SignalInput {
pub fn as_array(&self) -> &Array1<f32> {
&self.0
}
pub fn into_array(self) -> Array1<f32> {
self.0
}
}
impl AsRef<Array1<f32>> for SignalInput {
fn as_ref(&self) -> &Array1<f32> {
&self.0
}
}
#[derive(Default)]
pub struct KizzasiBuilder {
config: KizzasiConfig,
#[cfg(feature = "logic")]
guardrails: Option<GuardrailSet>,
}
impl KizzasiBuilder {
pub fn new() -> Self {
Self::default()
}
pub fn model_type(mut self, model_type: ModelType) -> Self {
self.config = self.config.model_type(model_type);
self
}
pub fn context_window(mut self, size: usize) -> Self {
self.config = self.config.context_window(size);
self
}
pub fn hidden_dim(mut self, dim: usize) -> Self {
self.config = self.config.hidden_dim(dim);
self
}
pub fn state_dim(mut self, dim: usize) -> Self {
self.config = self.config.state_dim(dim);
self
}
pub fn num_layers(mut self, n: usize) -> Self {
self.config = self.config.num_layers(n);
self
}
pub fn input_dim(mut self, dim: usize) -> Self {
self.config = self.config.input_dim(dim);
self
}
pub fn output_dim(mut self, dim: usize) -> Self {
self.config = self.config.output_dim(dim);
self
}
pub fn weights_path(mut self, path: &str) -> Self {
self.config = self.config.load_weights(path);
self
}
#[cfg(feature = "logic")]
pub fn guardrails(mut self, guardrails: GuardrailSet) -> Self {
self.guardrails = Some(guardrails);
self
}
pub fn build(self) -> KizzasiResult<Kizzasi> {
if self.config.get_input_dim() == 0 {
return Err(KizzasiError::Config("input_dim must be > 0".into()));
}
if self.config.get_output_dim() == 0 {
return Err(KizzasiError::Config("output_dim must be > 0".into()));
}
if self.config.get_hidden_dim() == 0 {
return Err(KizzasiError::Config("hidden_dim must be > 0".into()));
}
let mut predictor = Kizzasi::new(self.config)?;
#[cfg(feature = "logic")]
if let Some(guardrails) = self.guardrails {
predictor.set_guardrails(guardrails);
}
Ok(predictor)
}
}
impl KizzasiBuilder {
pub fn audio_preset() -> Self {
Self::new()
.model_type(ModelType::Mamba2)
.input_dim(1)
.output_dim(1)
.hidden_dim(256)
.state_dim(16)
.num_layers(4)
.context_window(8192)
}
pub fn robotics_preset(axes: usize) -> Self {
Self::new()
.model_type(ModelType::Mamba2)
.input_dim(axes)
.output_dim(axes)
.hidden_dim(128)
.state_dim(8)
.num_layers(3)
.context_window(1024)
}
pub fn sensor_preset(num_sensors: usize) -> Self {
Self::new()
.model_type(ModelType::Mamba2)
.input_dim(num_sensors)
.output_dim(num_sensors)
.hidden_dim(64)
.state_dim(8)
.num_layers(2)
.context_window(2048)
}
pub fn lightweight_preset(input_dim: usize, output_dim: usize) -> Self {
Self::new()
.model_type(ModelType::Mamba)
.input_dim(input_dim)
.output_dim(output_dim)
.hidden_dim(32)
.state_dim(4)
.num_layers(1)
.context_window(512)
}
pub fn video_preset(frame_features: usize) -> Self {
Self::new()
.model_type(ModelType::Mamba2)
.input_dim(frame_features)
.output_dim(frame_features)
.hidden_dim(512)
.state_dim(32)
.num_layers(6)
.context_window(16384) }
pub fn control_preset(state_dim_arg: usize, action_dim: usize) -> Self {
Self::new()
.model_type(ModelType::Mamba2)
.input_dim(state_dim_arg)
.output_dim(action_dim)
.hidden_dim(64)
.state_dim(8)
.num_layers(2)
.context_window(256) }
pub fn custom_preset() -> Self {
Self::new()
.model_type(ModelType::Mamba2)
.hidden_dim(128)
.state_dim(16)
.num_layers(3)
.context_window(4096)
}
}
pub struct Kizzasi {
ssm: SelectiveSSM,
config: KizzasiConfig,
#[cfg(feature = "logic")]
guardrails: Option<GuardrailSet>,
plugins: PluginManager,
}
impl Kizzasi {
pub fn new(config: KizzasiConfig) -> KizzasiResult<Self> {
let ssm = SelectiveSSM::new(config.clone())?;
Ok(Self {
ssm,
config,
#[cfg(feature = "logic")]
guardrails: None,
plugins: PluginManager::new(),
})
}
pub fn from_ssm(ssm: SelectiveSSM) -> KizzasiResult<Self> {
let config = ssm.config().clone();
Ok(Self {
ssm,
config,
#[cfg(feature = "logic")]
guardrails: None,
plugins: PluginManager::new(),
})
}
pub fn ssm(&self) -> &SelectiveSSM {
&self.ssm
}
pub fn ssm_mut(&mut self) -> &mut SelectiveSSM {
&mut self.ssm
}
pub fn add_plugin(&mut self, plugin: Box<dyn Plugin>) {
self.plugins.add_plugin(plugin);
}
pub fn remove_plugin(&mut self, name: &str) -> Option<Box<dyn Plugin>> {
self.plugins.remove_plugin(name)
}
pub fn plugins(&self) -> &PluginManager {
&self.plugins
}
pub fn plugins_mut(&mut self) -> &mut PluginManager {
&mut self.plugins
}
#[cfg(feature = "logic")]
pub fn set_guardrails(&mut self, guardrails: GuardrailSet) {
self.guardrails = Some(guardrails);
}
#[cfg(feature = "logic")]
pub fn clear_guardrails(&mut self) {
self.guardrails = None;
}
pub fn step(&mut self, input: &Array1<f32>) -> KizzasiResult<Array1<f32>> {
let input_dim = self.config.get_input_dim();
let output_dim = self.config.get_output_dim();
self.plugins
.execute_pre_process(input, input_dim, output_dim)?;
let transformed_input =
self.plugins
.transform_input(input.clone(), input_dim, output_dim)?;
let mut prediction = self.ssm.step(&transformed_input)?;
#[cfg(feature = "logic")]
if let Some(ref guardrails) = self.guardrails {
prediction = guardrails.constrain(&prediction)?;
}
let transformed_output = self
.plugins
.transform_output(prediction, input_dim, output_dim)?;
self.plugins
.execute_post_process(input, &transformed_output, input_dim, output_dim)?;
Ok(transformed_output)
}
pub fn step_slice(&mut self, input: &[f32]) -> KizzasiResult<Array1<f32>> {
let input_array = Array1::from_vec(input.to_vec());
self.step(&input_array)
}
pub fn step_inplace(&mut self, input: &[f32], output: &mut [f32]) -> KizzasiResult<()> {
let output_dim = self.config.get_output_dim();
if output.len() != output_dim {
return Err(KizzasiError::DimensionMismatch {
expected: output_dim,
actual: output.len(),
context: "Output buffer size must match output_dim".into(),
});
}
let result = self.step_slice(input)?;
let slice = result
.as_slice()
.ok_or_else(|| KizzasiError::inference("result array not contiguous"))?;
output.copy_from_slice(slice);
Ok(())
}
pub fn predict_n_inplace(
&mut self,
input: &Array1<f32>,
n_steps: usize,
output: &mut Array2<f32>,
) -> KizzasiResult<()> {
let output_dim = self.config.get_output_dim();
if output.shape() != [n_steps, output_dim] {
return Err(KizzasiError::DimensionMismatch {
expected: n_steps * output_dim,
actual: output.len(),
context: format!(
"Output buffer must be ({}, {}), got {:?}",
n_steps,
output_dim,
output.shape()
),
});
}
let mut current_input = input.clone();
for i in 0..n_steps {
let step_output = self.step(¤t_input)?;
for (j, &val) in step_output.iter().enumerate() {
output[[i, j]] = val;
}
current_input = step_output;
}
Ok(())
}
pub fn reset(&mut self) {
self.ssm.reset();
let input_dim = self.config.get_input_dim();
let output_dim = self.config.get_output_dim();
let _ = self.plugins.execute_on_reset(input_dim, output_dim);
}
pub fn context_window(&self) -> usize {
self.ssm.context_window()
}
pub fn config(&self) -> &KizzasiConfig {
&self.config
}
#[cfg(feature = "logic")]
pub fn has_guardrails(&self) -> bool {
self.guardrails.is_some()
}
#[cfg(feature = "logic")]
pub fn guardrails(&self) -> Option<&GuardrailSet> {
self.guardrails.as_ref()
}
#[cfg(feature = "logic")]
pub fn validate(&self, prediction: &Array1<f32>) -> bool {
if let Some(ref guardrails) = self.guardrails {
guardrails.validate(prediction)
} else {
true
}
}
#[cfg(feature = "logic")]
pub fn violation_loss(&self, prediction: &Array1<f32>) -> f32 {
if let Some(ref guardrails) = self.guardrails {
guardrails.violation_loss(prediction)
} else {
0.0
}
}
pub fn predict_n(&mut self, input: &Array1<f32>, n_steps: usize) -> KizzasiResult<Array2<f32>> {
let output_dim = self.config.get_output_dim();
let mut predictions = Array2::zeros((n_steps, output_dim));
let mut current_input = input.clone();
for i in 0..n_steps {
let output = self.step(¤t_input)?;
for (j, &val) in output.iter().enumerate() {
predictions[[i, j]] = val;
}
current_input = output;
}
Ok(predictions)
}
pub fn predict_until<F>(
&mut self,
input: &Array1<f32>,
max_steps: usize,
predicate: F,
) -> KizzasiResult<Vec<Array1<f32>>>
where
F: Fn(&Array1<f32>, usize) -> bool,
{
let mut predictions = Vec::with_capacity(max_steps);
let mut current_input = input.clone();
for step in 0..max_steps {
let output = self.step(¤t_input)?;
predictions.push(output.clone());
if predicate(&output, step) {
break;
}
current_input = output;
}
Ok(predictions)
}
pub fn predict_batch(&mut self, inputs: &[Array1<f32>]) -> KizzasiResult<Vec<Array1<f32>>> {
let mut outputs = Vec::with_capacity(inputs.len());
for input in inputs {
outputs.push(self.step(input)?);
}
Ok(outputs)
}
pub fn fork(&self) -> KizzasiResult<Self> {
let mut new_predictor = Self::new(self.config.clone())?;
#[cfg(feature = "logic")]
if let Some(ref guardrails) = self.guardrails {
new_predictor.guardrails = Some(guardrails.clone());
}
Ok(new_predictor)
}
pub fn hot_swap(
&mut self,
new_config: KizzasiConfig,
preserve_guardrails: bool,
) -> KizzasiResult<()> {
if new_config.get_input_dim() != self.config.get_input_dim() {
return Err(KizzasiError::DimensionMismatch {
expected: self.config.get_input_dim(),
actual: new_config.get_input_dim(),
context: "input_dim must match for hot-swap compatibility".into(),
});
}
if new_config.get_output_dim() != self.config.get_output_dim() {
return Err(KizzasiError::DimensionMismatch {
expected: self.config.get_output_dim(),
actual: new_config.get_output_dim(),
context: "output_dim must match for hot-swap compatibility".into(),
});
}
let new_ssm =
SelectiveSSM::new(new_config.clone()).map_err(|e| KizzasiError::ModelNotReady {
reason: format!("Failed to initialize new model: {}", e),
suggestion: "Check that the new configuration is valid and compatible".into(),
})?;
#[cfg(feature = "logic")]
let old_guardrails = if preserve_guardrails {
self.guardrails.clone()
} else {
None
};
self.ssm = new_ssm;
self.config = new_config;
#[cfg(feature = "logic")]
if preserve_guardrails {
self.guardrails = old_guardrails;
} else {
self.guardrails = None;
}
Ok(())
}
pub fn model_type(&self) -> ModelType {
self.config.get_model_type()
}
pub fn input_dim(&self) -> usize {
self.config.get_input_dim()
}
pub fn output_dim(&self) -> usize {
self.config.get_output_dim()
}
pub fn hidden_dim(&self) -> usize {
self.config.get_hidden_dim()
}
pub fn num_layers(&self) -> usize {
self.config.get_num_layers()
}
pub fn state_dim(&self) -> usize {
self.config.get_state_dim()
}
}
#[cfg(test)]
mod tests {
use super::*;
use kizzasi_core::ModelType;
#[test]
fn test_kizzasi_step() {
let config = KizzasiConfig::new()
.model_type(ModelType::Mamba2)
.input_dim(3)
.output_dim(3)
.hidden_dim(64)
.state_dim(8)
.num_layers(2);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let output = predictor.step(&input).unwrap();
assert_eq!(output.len(), 3);
}
#[test]
fn test_kizzasi_reset() {
let config = KizzasiConfig::new().input_dim(3).output_dim(3);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let _ = predictor.step(&input);
predictor.reset();
let output = predictor.step(&input).unwrap();
assert_eq!(output.len(), 3);
}
#[test]
fn test_predict_n() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(32)
.state_dim(4)
.num_layers(1);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let predictions = predictor.predict_n(&input, 5).unwrap();
assert_eq!(predictions.shape(), &[5, 3]);
}
#[test]
fn test_predict_until() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(32)
.state_dim(4)
.num_layers(1);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let predictions = predictor
.predict_until(&input, 10, |_, step| step >= 2)
.unwrap();
assert_eq!(predictions.len(), 3);
}
#[test]
fn test_predict_batch() {
let config = KizzasiConfig::new()
.input_dim(2)
.output_dim(2)
.hidden_dim(16)
.state_dim(4)
.num_layers(1);
let mut predictor = Kizzasi::new(config).unwrap();
let inputs = vec![
Array1::from_vec(vec![0.1, 0.2]),
Array1::from_vec(vec![0.3, 0.4]),
Array1::from_vec(vec![0.5, 0.6]),
];
let outputs = predictor.predict_batch(&inputs).unwrap();
assert_eq!(outputs.len(), 3);
for output in &outputs {
assert_eq!(output.len(), 2);
}
}
#[test]
fn test_fork() {
let config = KizzasiConfig::new()
.input_dim(2)
.output_dim(2)
.hidden_dim(16)
.state_dim(4)
.num_layers(1);
let predictor = Kizzasi::new(config).unwrap();
let forked = predictor.fork().unwrap();
assert_eq!(forked.context_window(), predictor.context_window());
}
#[test]
fn test_kizzasi_builder() {
let predictor = KizzasiBuilder::new()
.model_type(ModelType::Mamba2)
.input_dim(3)
.output_dim(3)
.hidden_dim(64)
.state_dim(8)
.num_layers(2)
.build()
.unwrap();
assert_eq!(predictor.context_window(), 8192);
}
#[test]
fn test_audio_preset() {
let predictor = KizzasiBuilder::audio_preset().build().unwrap();
assert_eq!(predictor.config().get_input_dim(), 1);
assert_eq!(predictor.config().get_hidden_dim(), 256);
}
#[test]
fn test_robotics_preset() {
let predictor = KizzasiBuilder::robotics_preset(6).build().unwrap();
assert_eq!(predictor.config().get_input_dim(), 6);
assert_eq!(predictor.config().get_output_dim(), 6);
}
#[test]
fn test_builder_validation() {
let result = KizzasiBuilder::new().input_dim(0).output_dim(1).build();
assert!(result.is_err());
}
#[test]
fn test_video_preset() {
let predictor = KizzasiBuilder::video_preset(256).build().unwrap();
assert_eq!(predictor.config().get_input_dim(), 256);
assert_eq!(predictor.config().get_output_dim(), 256);
assert_eq!(predictor.config().get_hidden_dim(), 512);
}
#[test]
fn test_control_preset() {
let predictor = KizzasiBuilder::control_preset(8, 4).build().unwrap();
assert_eq!(predictor.config().get_input_dim(), 8);
assert_eq!(predictor.config().get_output_dim(), 4);
assert_eq!(predictor.context_window(), 256);
}
#[test]
fn test_custom_preset() {
let predictor = KizzasiBuilder::custom_preset()
.input_dim(10)
.output_dim(10)
.build()
.unwrap();
assert_eq!(predictor.config().get_input_dim(), 10);
assert_eq!(predictor.config().get_hidden_dim(), 128);
}
#[test]
fn test_signal_input_from_f32() {
let input: SignalInput = 0.5f32.into();
assert_eq!(input.as_array().len(), 1);
assert_eq!(input.as_array()[0], 0.5);
}
#[test]
fn test_signal_input_from_vec() {
let input: SignalInput = vec![0.1, 0.2, 0.3].into();
assert_eq!(input.as_array().len(), 3);
assert_eq!(input.as_array()[0], 0.1);
}
#[test]
fn test_signal_input_from_slice() {
let data: &[f32] = &[0.1f32, 0.2, 0.3];
let input: SignalInput = data.into();
assert_eq!(input.as_array().len(), 3);
}
#[test]
fn test_signal_input_from_array() {
let input: SignalInput = [0.1f32, 0.2, 0.3].into();
assert_eq!(input.as_array().len(), 3);
assert_eq!(input.as_array()[2], 0.3);
}
#[test]
fn test_signal_input_into_array() {
let input: SignalInput = vec![0.1, 0.2].into();
let array = input.into_array();
assert_eq!(array.len(), 2);
}
#[test]
fn test_hot_swap_success() {
let config = KizzasiConfig::new()
.model_type(ModelType::Mamba2)
.input_dim(3)
.output_dim(3)
.hidden_dim(64)
.state_dim(8)
.num_layers(2);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let _ = predictor.step(&input).unwrap();
let new_config = KizzasiConfig::new()
.model_type(ModelType::Mamba) .input_dim(3) .output_dim(3)
.hidden_dim(128) .state_dim(16) .num_layers(4);
let result = predictor.hot_swap(new_config, false);
assert!(result.is_ok());
let output = predictor.step(&input).unwrap();
assert_eq!(output.len(), 3);
assert_eq!(predictor.model_type(), ModelType::Mamba);
assert_eq!(predictor.hidden_dim(), 128);
}
#[test]
fn test_hot_swap_dimension_mismatch_input() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(64);
let mut predictor = Kizzasi::new(config).unwrap();
let new_config = KizzasiConfig::new()
.input_dim(5) .output_dim(3)
.hidden_dim(64);
let result = predictor.hot_swap(new_config, false);
assert!(result.is_err());
if let Err(KizzasiError::DimensionMismatch {
expected, actual, ..
}) = result
{
assert_eq!(expected, 3);
assert_eq!(actual, 5);
} else {
panic!("Expected DimensionMismatch error");
}
}
#[test]
fn test_hot_swap_dimension_mismatch_output() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(64);
let mut predictor = Kizzasi::new(config).unwrap();
let new_config = KizzasiConfig::new()
.input_dim(3)
.output_dim(5) .hidden_dim(64);
let result = predictor.hot_swap(new_config, false);
assert!(result.is_err());
}
#[test]
#[cfg(feature = "logic")]
fn test_hot_swap_preserve_guardrails() {
use kizzasi_logic::{ConstraintBuilder, Guardrail, GuardrailSet};
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(64);
let mut predictor = Kizzasi::new(config).unwrap();
let mut guardrails = GuardrailSet::new();
let constraint = ConstraintBuilder::new()
.name("test_constraint")
.greater_eq(-1.0)
.less_eq(1.0)
.build()
.unwrap();
guardrails.add_global(Guardrail::new(constraint, false));
predictor.set_guardrails(guardrails);
assert!(predictor.has_guardrails());
let new_config = KizzasiConfig::new()
.model_type(ModelType::Mamba)
.input_dim(3)
.output_dim(3)
.hidden_dim(128);
predictor.hot_swap(new_config, true).unwrap();
assert!(predictor.has_guardrails());
}
#[test]
#[cfg(feature = "logic")]
fn test_hot_swap_discard_guardrails() {
use kizzasi_logic::{ConstraintBuilder, Guardrail, GuardrailSet};
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(64);
let mut predictor = Kizzasi::new(config).unwrap();
let mut guardrails = GuardrailSet::new();
let constraint = ConstraintBuilder::new()
.name("test_constraint")
.greater_eq(-1.0)
.less_eq(1.0)
.build()
.unwrap();
guardrails.add_global(Guardrail::new(constraint, false));
predictor.set_guardrails(guardrails);
assert!(predictor.has_guardrails());
let new_config = KizzasiConfig::new()
.model_type(ModelType::Mamba)
.input_dim(3)
.output_dim(3)
.hidden_dim(128);
predictor.hot_swap(new_config, false).unwrap();
assert!(!predictor.has_guardrails());
}
#[test]
fn test_accessor_methods() {
let config = KizzasiConfig::new()
.model_type(ModelType::Mamba2)
.input_dim(5)
.output_dim(7)
.hidden_dim(128)
.state_dim(16)
.num_layers(4);
let predictor = Kizzasi::new(config).unwrap();
assert_eq!(predictor.model_type(), ModelType::Mamba2);
assert_eq!(predictor.input_dim(), 5);
assert_eq!(predictor.output_dim(), 7);
assert_eq!(predictor.hidden_dim(), 128);
assert_eq!(predictor.state_dim(), 16);
assert_eq!(predictor.num_layers(), 4);
}
#[test]
fn test_step_slice() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(32)
.state_dim(4)
.num_layers(1);
let mut predictor = Kizzasi::new(config).unwrap();
let input_data = vec![0.1, 0.2, 0.3];
let output = predictor.step_slice(&input_data).unwrap();
assert_eq!(output.len(), 3);
let input_array = [0.4f32, 0.5, 0.6];
let output2 = predictor.step_slice(&input_array).unwrap();
assert_eq!(output2.len(), 3);
}
#[test]
fn test_step_inplace() {
let config = KizzasiConfig::new()
.input_dim(4)
.output_dim(4)
.hidden_dim(16)
.state_dim(4)
.num_layers(1);
let mut predictor = Kizzasi::new(config).unwrap();
let input = vec![0.1, 0.2, 0.3, 0.4];
let mut output = vec![0.0; 4];
predictor.step_inplace(&input, &mut output).unwrap();
for &val in &output {
assert!(val.is_finite());
}
}
#[test]
fn test_step_inplace_wrong_size() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(16);
let mut predictor = Kizzasi::new(config).unwrap();
let input = vec![0.1, 0.2, 0.3];
let mut output = vec![0.0; 5];
let result = predictor.step_inplace(&input, &mut output);
assert!(result.is_err());
if let Err(KizzasiError::DimensionMismatch {
expected, actual, ..
}) = result
{
assert_eq!(expected, 3);
assert_eq!(actual, 5);
}
}
#[test]
fn test_predict_n_inplace() {
let config = KizzasiConfig::new()
.input_dim(2)
.output_dim(2)
.hidden_dim(16)
.state_dim(4)
.num_layers(1);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2]);
let n_steps = 5;
let mut output = Array2::zeros((n_steps, 2));
predictor
.predict_n_inplace(&input, n_steps, &mut output)
.unwrap();
assert_eq!(output.shape(), &[5, 2]);
for &val in output.iter() {
assert!(val.is_finite());
}
}
#[test]
fn test_predict_n_inplace_wrong_shape() {
let config = KizzasiConfig::new()
.input_dim(2)
.output_dim(2)
.hidden_dim(16);
let mut predictor = Kizzasi::new(config).unwrap();
let input = Array1::from_vec(vec![0.1, 0.2]);
let mut output = Array2::zeros((5, 3));
let result = predictor.predict_n_inplace(&input, 5, &mut output);
assert!(result.is_err());
}
#[test]
fn test_zero_copy_equivalence() {
let config = KizzasiConfig::new()
.input_dim(3)
.output_dim(3)
.hidden_dim(32)
.state_dim(4)
.num_layers(1);
let mut predictor1 = Kizzasi::new(config.clone()).unwrap();
let mut predictor2 = Kizzasi::new(config).unwrap();
let input_array = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let input_slice = vec![0.1, 0.2, 0.3];
let output1 = predictor1.step(&input_array).unwrap();
let output2 = predictor2.step_slice(&input_slice).unwrap();
assert_eq!(output1.len(), output2.len());
}
}