use crate::engine::{EngineConfig, InferenceEngine};
use crate::error::{InferenceError, InferenceResult};
use kizzasi_logic::{ConstrainedInference, GuardrailSet};
use kizzasi_model::AutoregressiveModel;
use kizzasi_tokenizer::SignalTokenizer;
use scirs2_core::ndarray::Array1;
use std::sync::Arc;
pub type PreprocessHook = Arc<dyn Fn(&Array1<f32>) -> InferenceResult<Array1<f32>> + Send + Sync>;
pub type PostprocessHook = Arc<dyn Fn(&Array1<f32>) -> InferenceResult<Array1<f32>> + Send + Sync>;
pub struct Pipeline {
engine: InferenceEngine,
tokenizer: Option<Box<dyn SignalTokenizer>>,
use_tokenizer: bool,
constraints_enabled: bool,
guardrails: Option<GuardrailSet>,
preprocess_hooks: Vec<PreprocessHook>,
postprocess_hooks: Vec<PostprocessHook>,
}
impl Pipeline {
pub fn forward(&mut self, input: &Array1<f32>) -> InferenceResult<Array1<f32>> {
let mut preprocessed = input.clone();
for hook in &self.preprocess_hooks {
preprocessed = hook(&preprocessed)?;
}
let tokenized = if self.use_tokenizer {
if let Some(tokenizer) = &self.tokenizer {
tokenizer
.encode(&preprocessed)
.map_err(|e| InferenceError::TokenizationError(e.to_string()))?
} else {
return Err(InferenceError::TokenizationError(
"Tokenizer enabled but not provided".to_string(),
));
}
} else {
preprocessed
};
let output = self.engine.step(&tokenized)?;
let constrained = if self.constraints_enabled {
self.apply_constraints(&output)?
} else {
output
};
let decoded = if self.use_tokenizer {
if let Some(tokenizer) = &self.tokenizer {
tokenizer
.decode(&constrained)
.map_err(|e| InferenceError::TokenizationError(e.to_string()))?
} else {
return Err(InferenceError::TokenizationError(
"Tokenizer enabled but not provided".to_string(),
));
}
} else {
constrained
};
let mut postprocessed = decoded;
for hook in &self.postprocess_hooks {
postprocessed = hook(&postprocessed)?;
}
Ok(postprocessed)
}
fn apply_constraints(&self, output: &Array1<f32>) -> InferenceResult<Array1<f32>> {
if let Some(ref guardrails) = self.guardrails {
guardrails
.constrain(output)
.map_err(|e| InferenceError::ConstraintError(e.to_string()))
} else {
Ok(output.clone())
}
}
pub fn set_guardrails(&mut self, guardrails: GuardrailSet) {
self.guardrails = Some(guardrails);
self.constraints_enabled = true;
}
pub fn clear_guardrails(&mut self) {
self.guardrails = None;
self.constraints_enabled = false;
}
pub fn guardrails(&self) -> Option<&GuardrailSet> {
self.guardrails.as_ref()
}
pub fn rollout(
&mut self,
initial: &Array1<f32>,
steps: usize,
) -> InferenceResult<Vec<Array1<f32>>> {
let mut outputs = Vec::with_capacity(steps);
let mut current = initial.clone();
for _ in 0..steps {
let output = self.forward(¤t)?;
outputs.push(output.clone());
current = output;
}
Ok(outputs)
}
pub fn reset(&mut self) {
self.engine.reset();
}
pub fn engine(&self) -> &InferenceEngine {
&self.engine
}
pub fn engine_mut(&mut self) -> &mut InferenceEngine {
&mut self.engine
}
pub fn has_constraints(&self) -> bool {
self.constraints_enabled
}
pub fn has_tokenizer(&self) -> bool {
self.use_tokenizer && self.tokenizer.is_some()
}
pub fn add_preprocess_hook(&mut self, hook: PreprocessHook) {
self.preprocess_hooks.push(hook);
}
pub fn add_postprocess_hook(&mut self, hook: PostprocessHook) {
self.postprocess_hooks.push(hook);
}
pub fn num_preprocess_hooks(&self) -> usize {
self.preprocess_hooks.len()
}
pub fn num_postprocess_hooks(&self) -> usize {
self.postprocess_hooks.len()
}
pub fn clear_preprocess_hooks(&mut self) {
self.preprocess_hooks.clear();
}
pub fn clear_postprocess_hooks(&mut self) {
self.postprocess_hooks.clear();
}
}
pub struct PipelineBuilder {
engine_config: Option<EngineConfig>,
model: Option<Box<dyn AutoregressiveModel>>,
tokenizer: Option<Box<dyn SignalTokenizer>>,
use_tokenizer: bool,
constraints_enabled: bool,
guardrails: Option<GuardrailSet>,
preprocess_hooks: Vec<PreprocessHook>,
postprocess_hooks: Vec<PostprocessHook>,
}
impl PipelineBuilder {
pub fn new() -> Self {
Self {
engine_config: None,
model: None,
tokenizer: None,
use_tokenizer: false,
constraints_enabled: false,
guardrails: None,
preprocess_hooks: Vec::new(),
postprocess_hooks: Vec::new(),
}
}
pub fn engine_config(mut self, config: EngineConfig) -> Self {
self.engine_config = Some(config);
self
}
pub fn model(mut self, model: Box<dyn AutoregressiveModel>) -> Self {
self.model = Some(model);
self
}
pub fn tokenizer(mut self, tokenizer: Box<dyn SignalTokenizer>) -> Self {
self.tokenizer = Some(tokenizer);
self.use_tokenizer = true;
self
}
pub fn use_tokenizer(mut self, use_tok: bool) -> Self {
self.use_tokenizer = use_tok;
self
}
pub fn with_constraints(mut self) -> Self {
self.constraints_enabled = true;
self
}
pub fn guardrails(mut self, guardrails: GuardrailSet) -> Self {
self.guardrails = Some(guardrails);
self.constraints_enabled = true;
self
}
pub fn add_preprocess_hook(mut self, hook: PreprocessHook) -> Self {
self.preprocess_hooks.push(hook);
self
}
pub fn add_postprocess_hook(mut self, hook: PostprocessHook) -> Self {
self.postprocess_hooks.push(hook);
self
}
pub fn build(self) -> InferenceResult<Pipeline> {
let engine_config = self
.engine_config
.ok_or_else(|| InferenceError::PipelineConfig("engine_config not set".into()))?;
let engine = if let Some(model) = self.model {
InferenceEngine::with_model(engine_config, model)
} else {
InferenceEngine::new(engine_config)
};
Ok(Pipeline {
engine,
tokenizer: self.tokenizer,
use_tokenizer: self.use_tokenizer,
constraints_enabled: self.constraints_enabled,
guardrails: self.guardrails,
preprocess_hooks: self.preprocess_hooks,
postprocess_hooks: self.postprocess_hooks,
})
}
}
impl Default for PipelineBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sampling::SamplingConfig;
#[test]
fn test_pipeline_builder_basic() {
let engine_config = EngineConfig::new(3, 3);
let pipeline = PipelineBuilder::new()
.engine_config(engine_config)
.with_constraints()
.build();
assert!(pipeline.is_ok());
let p = pipeline.unwrap();
assert!(p.has_constraints());
assert!(!p.has_tokenizer());
}
#[test]
fn test_pipeline_missing_config() {
let result = PipelineBuilder::new().build();
assert!(result.is_err());
}
#[test]
fn test_pipeline_with_model() {
use kizzasi_model::s4::{S4Config, S4D};
let model_config = S4Config::new()
.input_dim(1)
.hidden_dim(64)
.state_dim(16)
.num_layers(2)
.diagonal(true);
let model = S4D::new(model_config).unwrap();
let engine_config = EngineConfig::new(1, 10);
let mut pipeline = PipelineBuilder::new()
.engine_config(engine_config)
.model(Box::new(model))
.build()
.unwrap();
let input = Array1::from_vec(vec![0.5]);
let output = pipeline.forward(&input);
assert!(output.is_ok());
}
#[test]
fn test_pipeline_rollout() {
use kizzasi_model::rwkv::{Rwkv, RwkvConfig};
let model_config = RwkvConfig::new()
.input_dim(1)
.hidden_dim(64)
.intermediate_dim(256)
.num_layers(2);
let model = Rwkv::new(model_config).unwrap();
let engine_config = EngineConfig::new(1, 10);
let mut pipeline = PipelineBuilder::new()
.engine_config(engine_config)
.model(Box::new(model))
.build()
.unwrap();
let initial = Array1::from_vec(vec![0.5]);
let outputs = pipeline.rollout(&initial, 5);
assert!(outputs.is_ok());
assert_eq!(outputs.unwrap().len(), 5);
}
#[test]
fn test_pipeline_reset() {
let engine_config = EngineConfig::new(1, 1);
let mut pipeline = PipelineBuilder::new()
.engine_config(engine_config)
.build()
.unwrap();
pipeline.reset();
assert_eq!(pipeline.engine().step_count(), 0);
}
#[test]
fn test_pipeline_with_sampling() {
use crate::sampling::SamplingStrategy;
use kizzasi_model::s4::{S4Config, S4D};
let model_config = S4Config::new()
.input_dim(1)
.hidden_dim(64)
.state_dim(16)
.num_layers(2)
.diagonal(true);
let model = S4D::new(model_config).unwrap();
let sampling = SamplingConfig::new()
.strategy(SamplingStrategy::TopK)
.top_k(5);
let engine_config = EngineConfig::new(1, 10)
.sampling(sampling)
.use_embeddings(true);
let mut pipeline = PipelineBuilder::new()
.engine_config(engine_config)
.model(Box::new(model))
.build()
.unwrap();
let input = Array1::from_vec(vec![0.5]);
let output = pipeline.forward(&input);
assert!(output.is_ok());
}
}