use crate::types::*;
use crate::models_v2::llama::{LlamaModelV2, LlamaConfig};
use crate::model_core::{Model, GenerationConfig, ModelInputs, ModelConfig};
use crate::tokenizer::Tokenizer;
use crate::tensor_core::{Tensor, Device};
pub struct Sampler {
_tensor_ops: crate::tensor_core::CpuTensorOpsImpl,
}
impl Sampler {
pub fn new() -> Self {
Self {
_tensor_ops: crate::tensor_core::CpuTensorOpsImpl::new(),
}
}
pub fn sample_greedy(&self, logits: &[f32]) -> ModelResult<u32> {
if logits.is_empty() {
return Err(ModelError::ComputationFailed("Empty logits".to_string()));
}
let mut max_idx = 0;
let mut max_val = logits[0];
for (idx, &val) in logits.iter().enumerate() {
if val > max_val {
max_val = val;
max_idx = idx;
}
}
Ok(max_idx as u32)
}
pub fn sample_temperature(&self, logits: &[f32], temperature: f32) -> ModelResult<u32> {
if logits.is_empty() {
return Err(ModelError::ComputationFailed("Empty logits".to_string()));
}
let scaled_logits: Vec<f32> = logits.iter().map(|&x| x / temperature).collect();
let probabilities = self.softmax(&scaled_logits)?;
self.sample_from_probabilities(&probabilities)
}
pub fn sample_top_p(&self, logits: &[f32], top_p: f32, temperature: f32) -> ModelResult<u32> {
if logits.is_empty() {
return Err(ModelError::ComputationFailed("Empty logits".to_string()));
}
let scaled_logits: Vec<f32> = logits.iter().map(|&x| x / temperature).collect();
let probabilities = self.softmax(&scaled_logits)?;
let mut indexed_probs: Vec<(usize, f32)> = probabilities
.iter()
.enumerate()
.map(|(i, &p)| (i, p))
.collect();
indexed_probs.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
let mut cumulative_prob = 0.0;
let mut cutoff_idx = indexed_probs.len();
for (i, (_, prob)) in indexed_probs.iter().enumerate() {
cumulative_prob += prob;
if cumulative_prob >= top_p {
cutoff_idx = i + 1;
break;
}
}
let mut filtered_probs = vec![0.0; logits.len()];
let mut total_prob = 0.0;
for (idx, prob) in indexed_probs.iter().take(cutoff_idx) {
filtered_probs[*idx] = *prob;
total_prob += prob;
}
if total_prob > 0.0 {
for prob in &mut filtered_probs {
*prob /= total_prob;
}
}
self.sample_from_probabilities(&filtered_probs)
}
fn softmax(&self, logits: &[f32]) -> ModelResult<Vec<f32>> {
if logits.is_empty() {
return Ok(Vec::new());
}
let max_val = logits.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut exp_vals = Vec::with_capacity(logits.len());
let mut sum = 0.0;
for &logit in logits {
let exp_val = (logit - max_val).exp();
exp_vals.push(exp_val);
sum += exp_val;
}
for exp_val in &mut exp_vals {
*exp_val /= sum;
}
Ok(exp_vals)
}
fn sample_from_probabilities(&self, probabilities: &[f32]) -> ModelResult<u32> {
use rand::Rng;
let mut rng = rand::thread_rng();
let random_val: f32 = rng.gen();
let mut cumulative = 0.0;
for (idx, &prob) in probabilities.iter().enumerate() {
cumulative += prob;
if random_val <= cumulative {
return Ok(idx as u32);
}
}
Ok((probabilities.len() - 1) as u32)
}
}
pub struct InferencePipeline {
model: LlamaModelV2,
tokenizer: Tokenizer,
sampler: Sampler,
}
impl InferencePipeline {
pub fn new(model: LlamaModelV2, tokenizer: Tokenizer) -> Self {
Self {
model,
tokenizer,
sampler: Sampler::new(),
}
}
pub fn generate(&self, prompt: &str, config: &GenerationConfig) -> ModelResult<String> {
let input_tokens = self.tokenizer.encode(prompt);
let generated_tokens = self.generate_tokens(&input_tokens, config)?;
let output_text = self.tokenizer.decode(&generated_tokens);
Ok(output_text)
}
fn generate_tokens(&self, input_tokens: &[u32], config: &GenerationConfig) -> ModelResult<Vec<u32>> {
let mut current_tokens = input_tokens.to_vec();
let mut generated_count = 0;
while generated_count < config.max_new_tokens {
let input_tensor = self.create_input_tensor(¤t_tokens)?;
let inputs = ModelInputs::Text {
input_ids: input_tensor,
attention_mask: None,
position_ids: None
};
let outputs = self.model.forward(&inputs).map_err(|e| ModelError::ComputationFailed(format!("Forward pass failed: {}", e)))?;
let logits = match outputs {
crate::model_core::ModelOutputs::Logits { logits, .. } => logits,
_ => return Err(ModelError::ComputationFailed("Expected logits output".to_string())),
};
let last_token_logits = self.extract_last_token_logits(&logits)?;
let next_token = if config.do_sample {
if config.top_p < 1.0 {
self.sampler.sample_top_p(&last_token_logits, config.top_p, config.temperature)?
} else {
self.sampler.sample_temperature(&last_token_logits, config.temperature)?
}
} else {
self.sampler.sample_greedy(&last_token_logits)?
};
if next_token == config.eos_token_id {
break;
}
current_tokens.push(next_token);
generated_count += 1;
}
Ok(current_tokens)
}
fn create_input_tensor(&self, tokens: &[u32]) -> ModelResult<Tensor> {
let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
let shape = &[1, tokens.len()];
Tensor::from_i64_slice(&tokens_i64, shape, &Device::CPU)
.map_err(|e| ModelError::ComputationFailed(format!("Failed to create input tensor: {}", e)))
}
fn extract_last_token_logits(&self, logits: &Tensor) -> ModelResult<Vec<f32>> {
let shape = logits.shape();
if shape.len() < 2 {
return Err(ModelError::ComputationFailed(
format!("Logits tensor has unexpected shape: {:?}", shape)
));
}
let candle_logits = logits.to_candle()
.map_err(|e| ModelError::ComputationFailed(format!("Failed to convert logits: {}", e)))?;
let vocab_size = if shape.len() == 3 {
shape[2]
} else if shape.len() == 2 {
shape[1]
} else {
return Err(ModelError::ComputationFailed(
format!("Unexpected logits shape: {:?}", shape)
));
};
let last_logits = if shape.len() == 3 {
let seq_len = shape[1];
candle_logits
.narrow(1, seq_len - 1, 1)
.map_err(|e| ModelError::ComputationFailed(format!("Failed to narrow: {}", e)))?
.squeeze(1)
.map_err(|e| ModelError::ComputationFailed(format!("Failed to squeeze: {}", e)))?
.squeeze(0)
.map_err(|e| ModelError::ComputationFailed(format!("Failed to squeeze batch: {}", e)))?
} else {
let seq_len = shape[0];
candle_logits
.narrow(0, seq_len - 1, 1)
.map_err(|e| ModelError::ComputationFailed(format!("Failed to narrow: {}", e)))?
.squeeze(0)
.map_err(|e| ModelError::ComputationFailed(format!("Failed to squeeze: {}", e)))?
};
let logits_vec: Vec<f32> = last_logits
.to_vec1()
.map_err(|e| ModelError::ComputationFailed(format!("Failed to convert to vec: {}", e)))?;
if logits_vec.len() != vocab_size {
return Err(ModelError::ComputationFailed(
format!("Logits size mismatch: got {}, expected {}", logits_vec.len(), vocab_size)
));
}
Ok(logits_vec)
}
pub fn model_config(&self) -> &LlamaConfig {
self.model.config()
}
pub fn tokenizer(&self) -> &Tokenizer {
&self.tokenizer
}
}
pub struct InferencePipelineBuilder {
model_config: Option<LlamaConfig>,
tokenizer: Option<Tokenizer>,
}
impl InferencePipelineBuilder {
pub fn new() -> Self {
Self {
model_config: None,
tokenizer: None,
}
}
pub fn with_model_config(mut self, config: LlamaConfig) -> Self {
self.model_config = Some(config);
self
}
pub fn with_tokenizer(mut self, tokenizer: Tokenizer) -> Self {
self.tokenizer = Some(tokenizer);
self
}
pub fn build(self) -> ModelResult<InferencePipeline> {
let model_config = self.model_config.ok_or_else(|| {
ModelError::InitializationFailed("Model config not provided".to_string())
})?;
let tokenizer = self.tokenizer.unwrap_or_else(Tokenizer::new);
let model = LlamaModelV2::new(model_config)
.map_err(|e| ModelError::InitializationFailed(format!("Failed to create model: {}", e)))?;
Ok(InferencePipeline::new(model, tokenizer))
}
}
#[derive(Debug, Clone)]
pub struct InferenceMetrics {
pub prompt_tokens: usize,
pub generated_tokens: usize,
pub total_tokens: usize,
pub inference_time_ms: f64,
pub tokens_per_second: f64,
}
impl InferenceMetrics {
pub fn new(prompt_tokens: usize, generated_tokens: usize, inference_time_ms: f64) -> Self {
let total_tokens = prompt_tokens + generated_tokens;
let tokens_per_second = if inference_time_ms > 0.0 {
(total_tokens as f64) / (inference_time_ms / 1000.0)
} else {
0.0
};
Self {
prompt_tokens,
generated_tokens,
total_tokens,
inference_time_ms,
tokens_per_second,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generation_config() {
let config = GenerationConfig::default();
assert_eq!(config.max_new_tokens, 100);
assert_eq!(config.temperature, 1.0);
assert_eq!(config.top_p, 0.9);
assert!(config.do_sample); }
#[test]
fn test_sampler_greedy() {
let sampler = Sampler::new();
let logits = vec![0.1, 0.8, 0.3, 0.5];
let token = sampler.sample_greedy(&logits).unwrap();
assert_eq!(token, 1); }
#[test]
fn test_sampler_empty_logits() {
let sampler = Sampler::new();
let logits = vec![];
let result = sampler.sample_greedy(&logits);
assert!(result.is_err());
}
#[test]
fn test_sampler_temperature() {
let sampler = Sampler::new();
let logits = vec![1.0, 2.0, 1.5];
let token = sampler.sample_temperature(&logits, 1.0).unwrap();
assert!(token < 3); }
#[test]
fn test_sampler_top_p() {
let sampler = Sampler::new();
let logits = vec![1.0, 3.0, 2.0, 0.5];
let token = sampler.sample_top_p(&logits, 0.8, 1.0).unwrap();
assert!(token < 4); }
#[test]
fn test_softmax() {
let sampler = Sampler::new();
let logits = vec![1.0, 2.0, 3.0];
let probs = sampler.softmax(&logits).unwrap();
assert_eq!(probs.len(), 3);
let sum: f32 = probs.iter().sum();
assert!((sum - 1.0).abs() < 1e-6);
assert!(probs[0] < probs[1]);
assert!(probs[1] < probs[2]);
}
#[test]
fn test_inference_pipeline_builder() {
let config = LlamaConfig {
vocab_size: 1000,
hidden_size: 128,
num_hidden_layers: 2,
num_attention_heads: 8,
..Default::default()
};
let tokenizer = Tokenizer::new();
let pipeline = InferencePipelineBuilder::new()
.with_model_config(config.clone())
.with_tokenizer(tokenizer)
.build()
.unwrap();
assert_eq!(pipeline.model_config().vocab_size(), config.vocab_size());
assert_eq!(pipeline.model_config().hidden_size(), config.hidden_size());
}
#[test]
fn test_inference_pipeline_generation() {
let config = LlamaConfig {
vocab_size: 1000,
hidden_size: 64,
num_hidden_layers: 1,
num_attention_heads: 4,
intermediate_size: 128,
max_position_embeddings: 64,
..Default::default()
};
let pipeline = InferencePipelineBuilder::new()
.with_model_config(config)
.build()
.unwrap();
let gen_config = GenerationConfig {
max_new_tokens: 5,
temperature: 1.0,
do_sample: false, ..Default::default()
};
let result = pipeline.generate("hello world", &gen_config);
match result {
Ok(output) => {
assert!(!output.is_empty());
println!("Generated: {}", output);
}
Err(e) => {
assert!(e.to_string().contains("matmul") || e.to_string().contains("Forward pass failed"));
println!("Expected error with dummy tensors: {}", e);
}
}
}
#[test]
fn test_inference_metrics() {
let metrics = InferenceMetrics::new(10, 20, 1000.0);
assert_eq!(metrics.prompt_tokens, 10);
assert_eq!(metrics.generated_tokens, 20);
assert_eq!(metrics.total_tokens, 30);
assert_eq!(metrics.inference_time_ms, 1000.0);
assert_eq!(metrics.tokens_per_second, 30.0);
}
#[test]
fn test_empty_prompt() {
let config = LlamaConfig {
vocab_size: 500,
hidden_size: 32,
num_hidden_layers: 1,
num_attention_heads: 2,
intermediate_size: 64,
max_position_embeddings: 32,
..Default::default()
};
let pipeline = InferencePipelineBuilder::new()
.with_model_config(config)
.build()
.unwrap();
let gen_config = GenerationConfig {
max_new_tokens: 3,
..Default::default()
};
let result = pipeline.generate("", &gen_config);
match result {
Ok(output) => {
println!("Generated from empty prompt: '{}'", output);
}
Err(e) => {
assert!(e.to_string().contains("matmul") || e.to_string().contains("Forward pass failed"));
println!("Expected error with dummy tensors: {}", e);
}
}
}
#[test]
fn test_builder_missing_config() {
let result = InferencePipelineBuilder::new().build();
assert!(result.is_err());
}
}