use crate::error::{WhisperError, WhisperResult};
use crate::model::encoder::{FeedForward, LayerNorm};
use crate::model::lfm2::layer::LayerNormNoBias;
use crate::model::lfm2::rope::{RopeConfig, RotaryEmbedding};
use crate::model::moonshine::MoonshineDecoderBlock;
use crate::model::{AttentionType, ModelConfig, MultiHeadAttention, PositionalEncoding};
use trueno::Matrix;
#[cfg(feature = "realizar-inference")]
use crate::model::encoder::FusedFFN;
#[derive(Debug, Clone)]
pub struct LayerKVCache {
pub key: Vec<f32>,
pub value: Vec<f32>,
pub seq_len: usize,
pub d_model: usize,
pub max_len: usize,
}
impl LayerKVCache {
#[must_use]
pub fn new(d_model: usize, max_len: usize) -> Self {
Self {
key: Vec::with_capacity(max_len * d_model),
value: Vec::with_capacity(max_len * d_model),
seq_len: 0,
d_model,
max_len,
}
}
#[must_use]
pub fn new_preallocated(d_model: usize, max_len: usize) -> Self {
let capacity = max_len * d_model;
let mut key = vec![0.0_f32; capacity];
let mut value = vec![0.0_f32; capacity];
key.truncate(0);
value.truncate(0);
key.reserve(capacity);
value.reserve(capacity);
Self {
key,
value,
seq_len: 0,
d_model,
max_len,
}
}
#[must_use]
pub fn remaining_capacity(&self) -> usize {
self.max_len.saturating_sub(self.seq_len)
}
#[must_use]
pub fn is_full(&self) -> bool {
self.seq_len >= self.max_len
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.seq_len == 0
}
#[must_use]
pub fn len(&self) -> usize {
self.seq_len
}
pub fn append(&mut self, new_key: &[f32], new_value: &[f32]) -> WhisperResult<()> {
let new_len = new_key.len() / self.d_model;
if new_key.len() != new_value.len() {
return Err(WhisperError::Model(
"key and value must have same size".into(),
));
}
if new_key.len() % self.d_model != 0 {
return Err(WhisperError::Model(
"key size not divisible by d_model".into(),
));
}
if self.seq_len + new_len > self.max_len {
return Err(WhisperError::Model(format!(
"cache overflow: {} + {} > {}",
self.seq_len, new_len, self.max_len
)));
}
self.key.extend_from_slice(new_key);
self.value.extend_from_slice(new_value);
self.seq_len += new_len;
Ok(())
}
#[must_use]
pub fn get_key(&self) -> &[f32] {
&self.key
}
#[must_use]
pub fn get_value(&self) -> &[f32] {
&self.value
}
pub fn clear(&mut self) {
self.key.clear();
self.value.clear();
self.seq_len = 0;
}
pub fn reset(&mut self) {
self.key.truncate(0);
self.value.truncate(0);
self.seq_len = 0;
}
pub fn append_batch(
&mut self,
keys: &[f32],
values: &[f32],
batch_size: usize,
) -> WhisperResult<()> {
let expected_len = batch_size * self.d_model;
if keys.len() != expected_len || values.len() != expected_len {
return Err(WhisperError::Model(format!(
"batch size mismatch: expected {} elements, got keys={}, values={}",
expected_len,
keys.len(),
values.len()
)));
}
if self.seq_len + batch_size > self.max_len {
return Err(WhisperError::Model(format!(
"cache overflow: {} + {} > {}",
self.seq_len, batch_size, self.max_len
)));
}
self.key.extend_from_slice(keys);
self.value.extend_from_slice(values);
self.seq_len += batch_size;
Ok(())
}
#[must_use]
pub fn get_key_range(&self, start: usize, end: usize) -> Option<&[f32]> {
if end > self.seq_len || start > end {
return None;
}
let start_idx = start * self.d_model;
let end_idx = end * self.d_model;
Some(&self.key[start_idx..end_idx])
}
#[must_use]
pub fn get_value_range(&self, start: usize, end: usize) -> Option<&[f32]> {
if end > self.seq_len || start > end {
return None;
}
let start_idx = start * self.d_model;
let end_idx = end * self.d_model;
Some(&self.value[start_idx..end_idx])
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
(self.key.len() + self.value.len()) * core::mem::size_of::<f32>()
}
#[must_use]
pub fn capacity_bytes(&self) -> usize {
(self.key.capacity() + self.value.capacity()) * core::mem::size_of::<f32>()
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)] pub struct LayerKVCacheTransposed {
pub key: Vec<f32>,
pub value_transposed: Vec<f32>,
pub seq_len: usize,
pub d_model: usize,
pub max_len: usize,
}
#[allow(dead_code)] impl LayerKVCacheTransposed {
#[must_use]
pub fn new(d_model: usize, max_len: usize) -> Self {
Self {
key: Vec::with_capacity(max_len * d_model),
value_transposed: Vec::with_capacity(max_len * d_model),
seq_len: 0,
d_model,
max_len,
}
}
#[must_use]
pub fn new_preallocated(d_model: usize, max_len: usize) -> Self {
let capacity = max_len * d_model;
Self {
key: vec![0.0_f32; capacity],
value_transposed: vec![0.0_f32; capacity],
seq_len: 0,
d_model,
max_len,
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.seq_len == 0
}
#[must_use]
pub fn len(&self) -> usize {
self.seq_len
}
pub fn append(&mut self, new_key: &[f32], new_value: &[f32]) -> WhisperResult<()> {
let new_len = new_key.len() / self.d_model;
if new_key.len() != new_value.len() {
return Err(WhisperError::Model(
"key and value must have same size".into(),
));
}
if new_key.len() % self.d_model != 0 {
return Err(WhisperError::Model(
"key size not divisible by d_model".into(),
));
}
if self.seq_len + new_len > self.max_len {
return Err(WhisperError::Model(format!(
"cache overflow: {} + {} > {}",
self.seq_len, new_len, self.max_len
)));
}
self.key.extend_from_slice(new_key);
if self.seq_len == 0 {
let total_new = new_len * self.d_model;
self.value_transposed.reserve(total_new);
for d in 0..self.d_model {
for t in 0..new_len {
self.value_transposed.push(new_value[t * self.d_model + d]);
}
}
} else {
let old_len = self.seq_len;
let new_total = (old_len + new_len) * self.d_model;
let mut new_transposed = Vec::with_capacity(new_total);
for d in 0..self.d_model {
let old_start = d * old_len;
let old_end = old_start + old_len;
new_transposed.extend_from_slice(&self.value_transposed[old_start..old_end]);
for t in 0..new_len {
new_transposed.push(new_value[t * self.d_model + d]);
}
}
self.value_transposed = new_transposed;
}
self.seq_len += new_len;
Ok(())
}
#[must_use]
pub fn get_key(&self) -> &[f32] {
&self.key
}
#[must_use]
pub fn get_value_transposed(&self) -> &[f32] {
&self.value_transposed
}
#[must_use]
pub fn get_value_feature(&self, feature_idx: usize) -> Option<&[f32]> {
if feature_idx >= self.d_model {
return None;
}
let start = feature_idx * self.seq_len;
let end = start + self.seq_len;
Some(&self.value_transposed[start..end])
}
#[must_use]
pub fn apply_attention(&self, scores: &[f32], query_len: usize) -> Vec<f32> {
let mut output = vec![0.0_f32; query_len * self.d_model];
for d in 0..self.d_model {
let v_feature = &self.value_transposed[d * self.seq_len..(d + 1) * self.seq_len];
for q in 0..query_len {
let score_row = &scores[q * self.seq_len..(q + 1) * self.seq_len];
let mut sum = 0.0_f32;
for (s, v) in score_row.iter().zip(v_feature.iter()) {
sum += s * v;
}
output[q * self.d_model + d] = sum;
}
}
output
}
pub fn clear(&mut self) {
self.key.clear();
self.value_transposed.clear();
self.seq_len = 0;
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
(self.key.len() + self.value_transposed.len()) * core::mem::size_of::<f32>()
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)] pub struct CircularKVBuffer {
key_cache: Vec<f32>,
value_cache: Vec<f32>,
head: usize,
valid_len: usize,
window_size: usize,
d_model: usize,
}
#[allow(dead_code)] impl CircularKVBuffer {
#[must_use]
pub fn new(window_size: usize, d_model: usize) -> Self {
let capacity = window_size * d_model;
Self {
key_cache: vec![0.0_f32; capacity],
value_cache: vec![0.0_f32; capacity],
head: 0,
valid_len: 0,
window_size,
d_model,
}
}
pub fn append(&mut self, key: &[f32], value: &[f32]) {
debug_assert_eq!(key.len(), self.d_model, "key length must match d_model");
debug_assert_eq!(value.len(), self.d_model, "value length must match d_model");
let pos = self.head % self.window_size;
let start = pos * self.d_model;
let end = start + self.d_model;
self.key_cache[start..end].copy_from_slice(key);
self.value_cache[start..end].copy_from_slice(value);
self.head += 1;
if self.valid_len < self.window_size {
self.valid_len += 1;
}
}
pub fn append_batch(&mut self, keys: &[f32], values: &[f32], batch_size: usize) {
debug_assert_eq!(keys.len(), batch_size * self.d_model);
debug_assert_eq!(values.len(), batch_size * self.d_model);
for i in 0..batch_size {
let offset = i * self.d_model;
self.append(
&keys[offset..offset + self.d_model],
&values[offset..offset + self.d_model],
);
}
}
#[must_use]
pub fn get_keys_linear(&self) -> Vec<f32> {
self.get_linear_view(&self.key_cache)
}
#[must_use]
pub fn get_values_linear(&self) -> Vec<f32> {
self.get_linear_view(&self.value_cache)
}
fn get_linear_view(&self, cache: &[f32]) -> Vec<f32> {
if self.valid_len < self.window_size {
cache[..self.valid_len * self.d_model].to_vec()
} else {
let start_pos = self.head % self.window_size;
let mut result = Vec::with_capacity(self.window_size * self.d_model);
let first_start = start_pos * self.d_model;
result.extend_from_slice(&cache[first_start..]);
result.extend_from_slice(&cache[..first_start]);
result
}
}
#[must_use]
pub fn len(&self) -> usize {
self.valid_len
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.valid_len == 0
}
#[must_use]
pub fn is_full(&self) -> bool {
self.valid_len >= self.window_size
}
#[must_use]
pub fn window_size(&self) -> usize {
self.window_size
}
pub fn reset(&mut self) {
self.head = 0;
self.valid_len = 0;
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
(self.key_cache.len() + self.value_cache.len()) * core::mem::size_of::<f32>()
}
}
#[derive(Debug, Clone)]
pub struct DecoderKVCache {
pub self_attn_cache: Vec<LayerKVCache>,
pub cross_attn_cache: Vec<LayerKVCache>,
pub n_layers: usize,
pub d_model: usize,
pub max_len: usize,
pub cross_attn_cached: bool,
pub seq_position: usize,
}
impl DecoderKVCache {
#[must_use]
pub fn new(n_layers: usize, d_model: usize, max_len: usize) -> Self {
let self_attn_cache = (0..n_layers)
.map(|_| LayerKVCache::new(d_model, max_len))
.collect();
let cross_attn_cache = (0..n_layers)
.map(|_| LayerKVCache::new(d_model, max_len * 4)) .collect();
Self {
self_attn_cache,
cross_attn_cache,
n_layers,
d_model,
max_len,
cross_attn_cached: false,
seq_position: 0,
}
}
#[must_use]
pub fn new_gqa(n_layers: usize, kv_dim: usize, d_model: usize, max_len: usize) -> Self {
let self_attn_cache = (0..n_layers)
.map(|_| LayerKVCache::new(kv_dim, max_len))
.collect();
let cross_attn_cache = (0..n_layers)
.map(|_| LayerKVCache::new(kv_dim, max_len * 4))
.collect();
Self {
self_attn_cache,
cross_attn_cache,
n_layers,
d_model,
max_len,
cross_attn_cached: false,
seq_position: 0,
}
}
#[must_use]
pub fn seq_len(&self) -> usize {
self.self_attn_cache.first().map_or(0, LayerKVCache::len)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.seq_len() == 0
}
pub fn clear(&mut self) {
for cache in &mut self.self_attn_cache {
cache.clear();
}
for cache in &mut self.cross_attn_cache {
cache.clear();
}
self.cross_attn_cached = false;
self.seq_position = 0;
}
pub fn clear_self_attn(&mut self) {
for cache in &mut self.self_attn_cache {
cache.clear();
}
}
pub fn increment_seq_len(&mut self) {
if let Some(cache) = self.self_attn_cache.first_mut() {
let dummy = vec![0.0f32; self.d_model];
cache.key.extend(&dummy);
cache.value.extend(&dummy);
}
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
let self_attn: usize = self
.self_attn_cache
.iter()
.map(|c| (c.key.len() + c.value.len()) * 4)
.sum();
let cross_attn: usize = self
.cross_attn_cache
.iter()
.map(|c| (c.key.len() + c.value.len()) * 4)
.sum();
self_attn + cross_attn
}
}
#[derive(Debug, Clone)]
pub struct StreamingKVCache {
inner: DecoderKVCache,
window_size: usize,
context_overlap: usize,
total_tokens: usize,
slide_count: usize,
}
impl StreamingKVCache {
#[must_use]
pub fn new(
n_layers: usize,
d_model: usize,
window_size: usize,
context_overlap: usize,
) -> Self {
Self {
inner: DecoderKVCache::new(n_layers, d_model, window_size),
window_size,
context_overlap: context_overlap.min(window_size / 2), total_tokens: 0,
slide_count: 0,
}
}
#[must_use]
pub fn low_latency(n_layers: usize, d_model: usize) -> Self {
Self::new(n_layers, d_model, 64, 16)
}
#[must_use]
pub fn ultra_low_latency(n_layers: usize, d_model: usize) -> Self {
Self::new(n_layers, d_model, 32, 8)
}
#[must_use]
pub fn standard(n_layers: usize, d_model: usize) -> Self {
Self::new(n_layers, d_model, 448, 64)
}
#[must_use]
pub fn seq_len(&self) -> usize {
self.inner.seq_len()
}
#[must_use]
pub fn total_tokens(&self) -> usize {
self.total_tokens
}
#[must_use]
pub fn slide_count(&self) -> usize {
self.slide_count
}
#[must_use]
pub fn remaining_capacity(&self) -> usize {
self.window_size.saturating_sub(self.seq_len())
}
#[must_use]
pub fn will_slide(&self) -> bool {
self.seq_len() >= self.window_size
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
#[must_use]
pub fn window_size(&self) -> usize {
self.window_size
}
#[must_use]
pub fn context_overlap(&self) -> usize {
self.context_overlap
}
#[must_use]
pub fn inner(&self) -> &DecoderKVCache {
&self.inner
}
pub fn inner_mut(&mut self) -> &mut DecoderKVCache {
&mut self.inner
}
pub fn append_with_slide(
&mut self,
layer_idx: usize,
key: &[f32],
value: &[f32],
) -> WhisperResult<()> {
let new_len = key.len() / self.inner.d_model;
if self.seq_len() + new_len > self.window_size {
self.slide_window()?;
}
self.inner.self_attn_cache[layer_idx].append(key, value)?;
self.total_tokens += new_len;
Ok(())
}
pub fn slide_window(&mut self) -> WhisperResult<()> {
let keep_from = self.seq_len().saturating_sub(self.context_overlap);
for cache in &mut self.inner.self_attn_cache {
if let (Some(k_range), Some(v_range)) = (
cache.get_key_range(keep_from, cache.len()),
cache.get_value_range(keep_from, cache.len()),
) {
let new_keys = k_range.to_vec();
let new_values = v_range.to_vec();
cache.reset();
cache.key.extend_from_slice(&new_keys);
cache.value.extend_from_slice(&new_values);
cache.seq_len = self.context_overlap;
}
}
self.slide_count += 1;
Ok(())
}
pub fn reset(&mut self) {
for cache in &mut self.inner.self_attn_cache {
cache.reset();
}
for cache in &mut self.inner.cross_attn_cache {
cache.reset();
}
self.inner.cross_attn_cached = false;
}
pub fn full_reset(&mut self) {
self.reset();
self.total_tokens = 0;
self.slide_count = 0;
}
pub fn warm_up(&mut self, layer_idx: usize, keys: &[f32], values: &[f32]) -> WhisperResult<()> {
if layer_idx >= self.inner.n_layers {
return Err(WhisperError::Model(format!(
"layer index {} out of bounds (max {})",
layer_idx, self.inner.n_layers
)));
}
let n_tokens = keys.len() / self.inner.d_model;
let tokens_to_use = n_tokens.min(self.context_overlap);
if tokens_to_use > 0 {
let start_idx = (n_tokens - tokens_to_use) * self.inner.d_model;
self.inner.self_attn_cache[layer_idx]
.append(&keys[start_idx..], &values[start_idx..])?;
}
Ok(())
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
self.inner.memory_bytes()
}
#[must_use]
pub fn stats(&self) -> StreamingCacheStats {
StreamingCacheStats {
seq_len: self.seq_len(),
total_tokens: self.total_tokens,
slide_count: self.slide_count,
window_size: self.window_size,
context_overlap: self.context_overlap,
memory_bytes: self.memory_bytes(),
}
}
}
#[derive(Debug, Clone)]
pub struct StreamingCacheStats {
pub seq_len: usize,
pub total_tokens: usize,
pub slide_count: usize,
pub window_size: usize,
pub context_overlap: usize,
pub memory_bytes: usize,
}
impl StreamingCacheStats {
#[must_use]
pub fn utilization(&self) -> f32 {
if self.window_size == 0 {
0.0
} else {
self.seq_len as f32 / self.window_size as f32
}
}
#[must_use]
pub fn tokens_per_slide(&self) -> f32 {
if self.slide_count == 0 {
self.total_tokens as f32
} else {
self.total_tokens as f32 / self.slide_count as f32
}
}
}
#[cfg(feature = "realizar-inference")]
#[allow(dead_code)] pub struct PagedDecoderKVCache {
layer_caches: Vec<crate::realizar_inference::PagedKvCache>,
n_layers: usize,
d_model: usize,
num_heads: usize,
head_dim: usize,
block_size: usize,
total_pages: usize,
layer_seq_ids: std::collections::HashMap<
crate::realizar_inference::SeqId,
Vec<crate::realizar_inference::SeqId>,
>,
seq_lengths: std::collections::HashMap<crate::realizar_inference::SeqId, usize>,
}
#[cfg(feature = "realizar-inference")]
#[allow(dead_code)] impl PagedDecoderKVCache {
#[must_use]
pub fn new(config: &ModelConfig, total_pages: usize) -> Self {
let n_layers = config.n_text_layer as usize;
let num_heads = config.n_text_head as usize;
let d_model = config.n_text_state as usize;
let head_dim = d_model / num_heads;
let block_size = 16;
let layer_caches = (0..n_layers)
.map(|_| {
crate::realizar_inference::PagedKvCache::new(
total_pages,
block_size,
num_heads,
head_dim,
)
})
.collect();
Self {
layer_caches,
n_layers,
d_model,
num_heads,
head_dim,
block_size,
total_pages,
layer_seq_ids: std::collections::HashMap::new(),
seq_lengths: std::collections::HashMap::new(),
}
}
#[must_use]
pub fn num_layers(&self) -> usize {
self.n_layers
}
#[must_use]
pub fn total_pages(&self) -> usize {
self.total_pages
}
#[must_use]
pub fn used_pages(&self) -> usize {
self.layer_caches
.iter()
.map(|c| c.stats().used_pages as usize)
.sum()
}
#[must_use]
pub fn has_sequence(&self, seq_id: crate::realizar_inference::SeqId) -> bool {
self.seq_lengths.contains_key(&seq_id)
}
#[must_use]
pub fn seq_len(&self, seq_id: crate::realizar_inference::SeqId) -> usize {
self.seq_lengths.get(&seq_id).copied().unwrap_or(0)
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
let used = self.used_pages();
let page_size = self.block_size * self.num_heads * self.head_dim;
used * page_size * 2 * core::mem::size_of::<f32>()
}
pub fn allocate_sequence(
&mut self,
initial_tokens: usize,
) -> WhisperResult<crate::realizar_inference::SeqId> {
let mut layer_ids = Vec::with_capacity(self.n_layers);
for cache in &mut self.layer_caches {
match cache.allocate_sequence(initial_tokens) {
Ok(id) => layer_ids.push(id),
Err(e) => {
for (i, &id) in layer_ids.iter().enumerate() {
self.layer_caches[i].free_sequence(id);
}
return Err(WhisperError::Model(format!(
"PagedKvCache allocation failed: {e}"
)));
}
}
}
let external_id = layer_ids[0];
self.layer_seq_ids.insert(external_id, layer_ids);
self.seq_lengths.insert(external_id, initial_tokens);
Ok(external_id)
}
pub fn free_sequence(&mut self, seq_id: crate::realizar_inference::SeqId) -> WhisperResult<()> {
let layer_ids = self
.layer_seq_ids
.remove(&seq_id)
.ok_or_else(|| WhisperError::Model("Sequence not found".into()))?;
for (i, layer_id) in layer_ids.into_iter().enumerate() {
self.layer_caches[i].free_sequence(layer_id);
}
self.seq_lengths.remove(&seq_id);
Ok(())
}
pub fn append(
&mut self,
seq_id: crate::realizar_inference::SeqId,
layer: usize,
key: &[f32],
value: &[f32],
) -> WhisperResult<()> {
if layer >= self.n_layers {
return Err(WhisperError::Model(format!(
"Layer {layer} out of range (max {})",
self.n_layers
)));
}
let layer_seq_id = self
.layer_seq_ids
.get(&seq_id)
.ok_or_else(|| WhisperError::Model("Sequence not found".into()))?[layer];
let current_len = self.seq_len(seq_id);
let token_size = self.num_heads * self.head_dim;
let pages_needed = (current_len + 1).div_ceil(self.block_size);
let current_pages = current_len.div_ceil(self.block_size);
if pages_needed > current_pages {
self.layer_caches[layer]
.extend(layer_seq_id, 1)
.map_err(|e| WhisperError::Model(format!("PagedKvCache extend failed: {e}")))?;
}
let page = self.layer_caches[layer]
.get_page_mut(layer_seq_id, current_len)
.map_err(|e| WhisperError::Model(format!("PagedKvCache get_page_mut failed: {e}")))?;
let offset_in_page = (current_len % self.block_size) * token_size;
page.keys[offset_in_page..offset_in_page + key.len()].copy_from_slice(key);
page.values[offset_in_page..offset_in_page + value.len()].copy_from_slice(value);
page.num_tokens = (current_len % self.block_size) + 1;
Ok(())
}
pub fn increment_seq_len(&mut self, seq_id: crate::realizar_inference::SeqId) {
*self.seq_lengths.entry(seq_id).or_insert(0) += 1;
}
pub fn get_kv(
&self,
seq_id: crate::realizar_inference::SeqId,
layer: usize,
) -> WhisperResult<(Vec<f32>, Vec<f32>)> {
if layer >= self.n_layers {
return Err(WhisperError::Model(format!(
"Layer {layer} out of range (max {})",
self.n_layers
)));
}
let seq_len = self.seq_len(seq_id);
if seq_len == 0 {
return Ok((Vec::new(), Vec::new()));
}
self.get_all_kv(seq_id, layer).map(|(keys, values)| {
let token_size = self.num_heads * self.head_dim;
let start = (seq_len - 1) * token_size;
(
keys[start..start + token_size].to_vec(),
values[start..start + token_size].to_vec(),
)
})
}
pub fn get_all_kv(
&self,
seq_id: crate::realizar_inference::SeqId,
layer: usize,
) -> WhisperResult<(Vec<f32>, Vec<f32>)> {
if layer >= self.n_layers {
return Err(WhisperError::Model(format!(
"Layer {layer} out of range (max {})",
self.n_layers
)));
}
let layer_seq_id = self
.layer_seq_ids
.get(&seq_id)
.ok_or_else(|| WhisperError::Model("Sequence not found".into()))?[layer];
let seq_len = self.seq_len(seq_id);
if seq_len == 0 {
return Ok((Vec::new(), Vec::new()));
}
let token_size = self.num_heads * self.head_dim;
let mut all_keys = Vec::with_capacity(seq_len * token_size);
let mut all_values = Vec::with_capacity(seq_len * token_size);
for token_pos in 0..seq_len {
let page = self.layer_caches[layer]
.get_page(layer_seq_id, token_pos)
.map_err(|e| WhisperError::Model(format!("PagedKvCache get_page failed: {e}")))?;
let offset_in_page = (token_pos % self.block_size) * token_size;
all_keys.extend_from_slice(&page.keys[offset_in_page..offset_in_page + token_size]);
all_values.extend_from_slice(&page.values[offset_in_page..offset_in_page + token_size]);
}
Ok((all_keys, all_values))
}
pub fn get_all_kv_n(
&self,
seq_id: crate::realizar_inference::SeqId,
layer: usize,
n_tokens: usize,
) -> WhisperResult<(Vec<f32>, Vec<f32>)> {
if layer >= self.n_layers {
return Err(WhisperError::Model(format!(
"Layer {layer} out of range (max {})",
self.n_layers
)));
}
if n_tokens == 0 {
return Ok((Vec::new(), Vec::new()));
}
let layer_seq_id = self
.layer_seq_ids
.get(&seq_id)
.ok_or_else(|| WhisperError::Model("Sequence not found".into()))?[layer];
let token_size = self.num_heads * self.head_dim;
let mut all_keys = Vec::with_capacity(n_tokens * token_size);
let mut all_values = Vec::with_capacity(n_tokens * token_size);
for token_pos in 0..n_tokens {
let page = self.layer_caches[layer]
.get_page(layer_seq_id, token_pos)
.map_err(|e| WhisperError::Model(format!("PagedKvCache get_page failed: {e}")))?;
let offset_in_page = (token_pos % self.block_size) * token_size;
all_keys.extend_from_slice(&page.keys[offset_in_page..offset_in_page + token_size]);
all_values.extend_from_slice(&page.values[offset_in_page..offset_in_page + token_size]);
}
Ok((all_keys, all_values))
}
}
#[derive(Debug, Clone)]
pub struct BatchDecoderCache {
caches: Vec<DecoderKVCache>,
pub n_layers: usize,
pub d_model: usize,
pub max_len: usize,
}
impl BatchDecoderCache {
#[must_use]
pub fn new(batch_size: usize, n_layers: usize, d_model: usize, max_len: usize) -> Self {
let caches = (0..batch_size)
.map(|_| DecoderKVCache::new(n_layers, d_model, max_len))
.collect();
Self {
caches,
n_layers,
d_model,
max_len,
}
}
#[must_use]
pub fn batch_size(&self) -> usize {
self.caches.len()
}
#[must_use]
pub fn get_cache(&self, index: usize) -> Option<&DecoderKVCache> {
self.caches.get(index)
}
pub fn get_cache_mut(&mut self, index: usize) -> Option<&mut DecoderKVCache> {
self.caches.get_mut(index)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.caches.iter().all(DecoderKVCache::is_empty)
}
pub fn clear_all(&mut self) {
for cache in &mut self.caches {
cache.clear();
}
}
#[must_use]
pub fn seq_lengths(&self) -> Vec<usize> {
self.caches.iter().map(DecoderKVCache::seq_len).collect()
}
#[must_use]
pub fn max_seq_len(&self) -> usize {
self.caches
.iter()
.map(DecoderKVCache::seq_len)
.max()
.unwrap_or(0)
}
#[must_use]
pub fn memory_bytes(&self) -> usize {
self.caches.iter().map(DecoderKVCache::memory_bytes).sum()
}
pub fn fork(&mut self, parent_id: usize, new_id: usize) -> WhisperResult<()> {
if parent_id >= self.caches.len() || new_id >= self.caches.len() {
return Err(WhisperError::Model("invalid cache fork index".into()));
}
if parent_id != new_id {
let parent_cache = self.caches[parent_id].clone();
self.caches[new_id] = parent_cache;
}
Ok(())
}
pub fn prune(&mut self, id: usize) -> WhisperResult<()> {
if id >= self.caches.len() {
return Err(WhisperError::Model("invalid cache prune index".into()));
}
self.caches[id].clear();
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct BatchDecoderOutput {
pub logits: Vec<Vec<f32>>,
pub seq_lengths: Vec<usize>,
}
impl BatchDecoderOutput {
#[must_use]
pub fn batch_size(&self) -> usize {
self.logits.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.logits.is_empty()
}
#[must_use]
pub fn get_logits(&self, index: usize) -> Option<&Vec<f32>> {
self.logits.get(index)
}
}
pub struct DecoderScratch {
normed: Vec<f32>,
q: Vec<f32>,
k_new: Vec<f32>,
v_new: Vec<f32>,
sa_proj: Vec<f32>,
normed2: Vec<f32>,
cross_q: Vec<f32>,
cross_proj: Vec<f32>,
normed3: Vec<f32>,
ffn_hidden: Vec<f32>,
ffn_out: Vec<f32>,
ln_post_out: Vec<f32>,
logits: Vec<f32>,
qkv: Vec<f32>,
attn: AttentionScratch,
}
pub struct AttentionScratch {
k_head: Vec<f32>,
v_head: Vec<f32>,
output: Vec<f32>,
scores: Vec<f32>,
weights: Vec<f32>,
}
impl DecoderScratch {
#[must_use]
pub fn new_with_attn(
d_model: usize,
d_ff: usize,
d_head: usize,
max_len: usize,
n_vocab: usize,
) -> Self {
Self {
normed: vec![0.0; d_model],
q: vec![0.0; d_model],
k_new: vec![0.0; d_model],
v_new: vec![0.0; d_model],
sa_proj: vec![0.0; d_model],
normed2: vec![0.0; d_model],
cross_q: vec![0.0; d_model],
cross_proj: vec![0.0; d_model],
normed3: vec![0.0; d_model],
ffn_hidden: vec![0.0; d_ff],
ffn_out: vec![0.0; d_model],
ln_post_out: vec![0.0; d_model],
logits: vec![0.0; n_vocab],
qkv: vec![0.0; 3 * d_model],
attn: AttentionScratch {
k_head: vec![0.0; max_len * d_head],
v_head: vec![0.0; max_len * d_head],
output: vec![0.0; d_model],
scores: vec![0.0; max_len],
weights: vec![0.0; max_len],
},
}
}
}
#[derive(Debug, Clone)]
pub struct DecoderBlock {
pub self_attn: MultiHeadAttention,
pub ln1: LayerNorm,
pub cross_attn: MultiHeadAttention,
pub ln2: LayerNorm,
pub ffn: FeedForward,
pub ln3: LayerNorm,
}
impl DecoderBlock {
#[must_use]
pub fn new(d_model: usize, n_heads: usize, d_ff: usize) -> Self {
Self {
self_attn: MultiHeadAttention::new(n_heads, d_model),
ln1: LayerNorm::new(d_model),
cross_attn: MultiHeadAttention::new(n_heads, d_model),
ln2: LayerNorm::new(d_model),
ffn: FeedForward::new(d_model, d_ff),
ln3: LayerNorm::new(d_model),
}
}
pub fn forward(
&self,
x: &[f32],
encoder_output: &[f32],
causal_mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let normed = self.ln1.forward(x)?;
let attn_out = self.self_attn.forward(&normed, causal_mask)?;
let mut residual: Vec<f32> = x.iter().zip(attn_out.iter()).map(|(a, b)| a + b).collect();
let normed = self.ln2.forward(&residual)?;
let cross_out = self
.cross_attn
.forward_cross_dispatch(&normed, encoder_output, None)?;
for (r, c) in residual.iter_mut().zip(cross_out.iter()) {
*r += c;
}
let normed = self.ln3.forward(&residual)?;
let ffn_out = self.ffn.forward(&normed)?;
for (r, f) in residual.iter_mut().zip(ffn_out.iter()) {
*r += f;
}
Ok(residual)
}
#[cfg(feature = "realizar-inference")]
pub fn forward_fused(
&self,
x: &[f32],
encoder_output: &[f32],
causal_mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let normed = self.ln1.forward(x)?;
let attn_out = self.self_attn.forward(&normed, causal_mask)?;
let mut residual: Vec<f32> = x.iter().zip(attn_out.iter()).map(|(a, b)| a + b).collect();
let normed = self.ln2.forward(&residual)?;
let cross_out = self
.cross_attn
.forward_cross_dispatch(&normed, encoder_output, None)?;
for (r, c) in residual.iter_mut().zip(cross_out.iter()) {
*r += c;
}
let fused = self.create_fused_ffn()?;
let ffn_out = fused.forward(&residual)?;
for (r, f) in residual.iter_mut().zip(ffn_out.iter()) {
*r += f;
}
Ok(residual)
}
pub fn finalize_weights(&mut self) {
self.self_attn.finalize_weights();
self.self_attn.fuse_qkv_weights();
self.cross_attn.finalize_weights();
self.cross_attn.fuse_qkv_weights();
self.ffn.finalize_weights();
}
#[must_use]
pub fn is_finalized(&self) -> bool {
self.self_attn.is_finalized() && self.cross_attn.is_finalized() && self.ffn.is_finalized()
}
pub fn self_attn_mut(&mut self) -> &mut MultiHeadAttention {
&mut self.self_attn
}
pub fn cross_attn_mut(&mut self) -> &mut MultiHeadAttention {
&mut self.cross_attn
}
pub fn ffn_mut(&mut self) -> &mut FeedForward {
&mut self.ffn
}
pub fn ln1_mut(&mut self) -> &mut LayerNorm {
&mut self.ln1
}
pub fn ln2_mut(&mut self) -> &mut LayerNorm {
&mut self.ln2
}
pub fn ln3_mut(&mut self) -> &mut LayerNorm {
&mut self.ln3
}
#[cfg(feature = "realizar-inference")]
pub fn create_fused_ffn(&self) -> WhisperResult<FusedFFN> {
let d_model = self.ln3.weight.len();
let d_ff = self.ffn.fc1.bias.len();
let mut fused = FusedFFN::new(d_model, d_ff)?;
fused.set_norm_weights(&self.ln3.weight, &self.ln3.bias);
fused.set_fc1_weights(&self.ffn.fc1.weight, &self.ffn.fc1.bias);
fused.set_fc2_weights(&self.ffn.fc2.weight, &self.ffn.fc2.bias);
Ok(fused)
}
}
#[derive(Debug, Clone)]
pub struct Decoder {
n_layers: usize,
d_model: usize,
n_heads: usize,
blocks: Vec<DecoderBlock>,
moonshine_blocks: Vec<MoonshineDecoderBlock>,
rope: Option<RotaryEmbedding>,
ln_post: LayerNorm,
ln_post_rms: Option<LayerNormNoBias>,
token_embedding: Vec<f32>,
token_embedding_transposed: Matrix<f32>,
token_embedding_f16: Option<Vec<u16>>,
positional_embedding: Vec<f32>,
n_vocab: usize,
max_len: usize,
positional_encoding: PositionalEncoding,
attention_type: AttentionType,
}
impl Decoder {
#[must_use]
pub fn new(config: &ModelConfig) -> Self {
let n_layers = config.n_text_layer as usize;
let d_model = config.n_text_state as usize;
let n_heads = config.n_text_head as usize;
let d_ff = d_model * 4;
let n_vocab = config.n_vocab as usize;
let max_len = config.n_text_ctx as usize;
let (blocks, moon_blocks, rope, ln_post_rms) = match config.attention_type {
AttentionType::Mha => {
let blocks: Vec<DecoderBlock> = (0..n_layers)
.map(|_| DecoderBlock::new(d_model, n_heads, d_ff))
.collect();
(blocks, Vec::new(), None, None)
}
AttentionType::Gqa { kv_heads } => {
let head_dim = d_model / n_heads;
let intermediate_size = d_model * 4;
let mut moon_blocks = Vec::with_capacity(n_layers);
for _ in 0..n_layers {
match MoonshineDecoderBlock::new(
d_model,
n_heads,
kv_heads as usize,
intermediate_size,
) {
Ok(block) => moon_blocks.push(block),
Err(_) => return Self::fallback_decoder(config),
}
}
let rotary_dim = (head_dim as f64 * 0.9).floor() as usize;
let rotary_dim = rotary_dim - (rotary_dim % 2);
let padded_hd = head_dim.div_ceil(8) * 8;
let Ok(rope_emb) = RotaryEmbedding::new(RopeConfig {
head_dim: padded_hd,
base: 10000.0,
max_seq_len: 2048,
rotary_dim: Some(rotary_dim),
}) else {
return Self::fallback_decoder(config);
};
(
Vec::new(),
moon_blocks,
Some(rope_emb),
Some(LayerNormNoBias::new(d_model)),
)
}
};
let positional_embedding = vec![0.0_f32; max_len * d_model];
let token_embedding = vec![0.0_f32; n_vocab * d_model];
let transposed_data = crate::simd::transpose(&token_embedding, n_vocab, d_model);
let token_embedding_transposed = Matrix::from_vec(d_model, n_vocab, transposed_data)
.unwrap_or_else(|_| Matrix::zeros(d_model, n_vocab));
Self {
n_layers,
d_model,
n_heads,
blocks,
moonshine_blocks: moon_blocks,
rope,
ln_post: LayerNorm::new(d_model),
ln_post_rms,
token_embedding,
token_embedding_transposed,
token_embedding_f16: None,
positional_embedding,
n_vocab,
max_len,
positional_encoding: config.positional_encoding,
attention_type: config.attention_type,
}
}
fn fallback_decoder(config: &ModelConfig) -> Self {
let d_model = config.n_text_state as usize;
let n_vocab = config.n_vocab as usize;
let max_len = config.n_text_ctx as usize;
let token_embedding = vec![0.0_f32; n_vocab * d_model];
let transposed_data = crate::simd::transpose(&token_embedding, n_vocab, d_model);
let token_embedding_transposed = Matrix::from_vec(d_model, n_vocab, transposed_data)
.unwrap_or_else(|_| Matrix::zeros(d_model, n_vocab));
Self {
n_layers: 0,
d_model,
n_heads: config.n_text_head as usize,
blocks: Vec::new(),
moonshine_blocks: Vec::new(),
rope: None,
ln_post: LayerNorm::new(d_model),
ln_post_rms: None,
token_embedding,
token_embedding_transposed,
token_embedding_f16: None,
positional_embedding: vec![0.0_f32; max_len * d_model],
n_vocab,
max_len,
positional_encoding: config.positional_encoding,
attention_type: config.attention_type,
}
}
fn update_embedding_transpose(&mut self) {
let transposed_data =
crate::simd::transpose(&self.token_embedding, self.n_vocab, self.d_model);
self.token_embedding_transposed =
Matrix::from_vec(self.d_model, self.n_vocab, transposed_data)
.unwrap_or_else(|_| Matrix::zeros(self.d_model, self.n_vocab));
}
pub fn finalize_weights(&mut self) {
for block in &mut self.blocks {
block.finalize_weights();
}
self.update_embedding_transpose();
}
pub fn convert_embeddings_to_f16(&mut self) {
if self.token_embedding_f16.is_some() || self.token_embedding.is_empty() {
return;
}
self.token_embedding_f16 = Some(crate::simd::quant_f32_to_f16(&self.token_embedding));
}
pub fn convert_to_f16(&mut self) {
for block in &mut self.blocks {
block.self_attn.convert_to_f16();
block.cross_attn.convert_to_f16();
block.ffn.convert_to_f16();
}
self.convert_embeddings_to_f16();
}
#[must_use]
pub fn is_finalized(&self) -> bool {
if self.moonshine_blocks.is_empty() {
self.blocks.iter().all(DecoderBlock::is_finalized)
} else {
true
}
}
#[cfg(feature = "realizar-inference")]
pub fn initialize_fused_ffn(&mut self) -> WhisperResult<()> {
for (i, block) in self.blocks.iter().enumerate() {
block.create_fused_ffn().map_err(|e| {
WhisperError::Model(format!("Block {i} failed to create FusedFFN: {e}"))
})?;
}
Ok(())
}
#[allow(clippy::no_effect_underscore_binding)]
pub fn forward(&self, tokens: &[u32], encoder_output: &[f32]) -> WhisperResult<Vec<f32>> {
let _span = crate::trace_enter!("step_h_decode");
let seq_len = tokens.len();
if seq_len == 0 {
return Err(WhisperError::Model("empty token sequence".into()));
}
if seq_len > self.max_len {
return Err(WhisperError::Model(format!(
"sequence length {} exceeds max {}",
seq_len, self.max_len
)));
}
if encoder_output.len() % self.d_model != 0 {
return Err(WhisperError::Model("encoder output size mismatch".into()));
}
let mut x = self.embed_tokens(tokens)?;
let enc_seq_len = encoder_output.len() / self.d_model;
if self.rope.is_some() {
let rope = self
.rope
.as_ref()
.ok_or_else(|| WhisperError::Model("Moonshine decoder requires RoPE".into()))?;
for block in &self.moonshine_blocks {
x = block.forward(&x, encoder_output, seq_len, enc_seq_len, rope)?;
}
if let Some(ref rms) = self.ln_post_rms {
x = rms.forward(&x, seq_len)?;
}
} else {
for pos in 0..seq_len {
for d in 0..self.d_model {
x[pos * self.d_model + d] += self.positional_embedding[pos * self.d_model + d];
}
}
let causal_mask = MultiHeadAttention::causal_mask(seq_len);
for block in &self.blocks {
x = block.forward(&x, encoder_output, Some(&causal_mask))?;
}
x = self.ln_post.forward(&x)?;
}
Ok(self.project_to_vocab(&x, seq_len))
}
#[allow(clippy::no_effect_underscore_binding)]
pub fn forward_probed(
&self,
tokens: &[u32],
encoder_output: &[f32],
probe: &mut crate::probe::ActivationProbe,
) -> WhisperResult<Vec<f32>> {
let seq_len = tokens.len();
if seq_len == 0 {
return Err(WhisperError::Model("empty token sequence".into()));
}
if seq_len > self.max_len {
return Err(WhisperError::Model(format!(
"sequence length {} exceeds max {}",
seq_len, self.max_len
)));
}
if encoder_output.len() % self.d_model != 0 {
return Err(WhisperError::Model("encoder output size mismatch".into()));
}
let mut x = self.embed_tokens(tokens)?;
probe.record("decoder.token_emb", &x, &[seq_len, self.d_model]);
let enc_seq_len = encoder_output.len() / self.d_model;
if self.rope.is_some() {
let rope = self
.rope
.as_ref()
.ok_or_else(|| WhisperError::Model("Moonshine decoder requires RoPE".into()))?;
for (i, block) in self.moonshine_blocks.iter().enumerate() {
x = block.forward_probed(
&x,
encoder_output,
seq_len,
enc_seq_len,
rope,
i,
probe,
)?;
}
if let Some(ref rms) = self.ln_post_rms {
x = rms.forward(&x, seq_len)?;
}
} else {
for pos in 0..seq_len {
for d in 0..self.d_model {
x[pos * self.d_model + d] += self.positional_embedding[pos * self.d_model + d];
}
}
let causal_mask = MultiHeadAttention::causal_mask(seq_len);
for block in &self.blocks {
x = block.forward(&x, encoder_output, Some(&causal_mask))?;
}
x = self.ln_post.forward(&x)?;
}
probe.record("decoder.ln_post_out", &x, &[seq_len, self.d_model]);
let logits = self.project_to_vocab(&x, seq_len);
probe.record("decoder.logits", &logits, &[seq_len, self.n_vocab]);
Ok(logits)
}
#[allow(clippy::similar_names, clippy::type_complexity)]
pub fn forward_traced(
&self,
tokens: &[u32],
encoder_output: &[f32],
) -> WhisperResult<(Vec<f32>, Vec<(String, f32)>)> {
let seq_len = tokens.len();
let mut trace: Vec<(String, f32)> = Vec::new();
if seq_len == 0 {
return Err(WhisperError::Model("empty token sequence".into()));
}
if seq_len > self.max_len {
return Err(WhisperError::Model(format!(
"sequence length {} exceeds max {}",
seq_len, self.max_len
)));
}
if encoder_output.len() % self.d_model != 0 {
return Err(WhisperError::Model("encoder output size mismatch".into()));
}
let mut x = self.embed_tokens(tokens)?;
let l2_token_emb: f32 = x.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("token_emb".to_string(), l2_token_emb));
if self.positional_encoding == PositionalEncoding::Sinusoidal {
for pos in 0..seq_len {
for d in 0..self.d_model {
x[pos * self.d_model + d] += self.positional_embedding[pos * self.d_model + d];
}
}
}
let l2_pos_emb: f32 = x.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("after_pos_emb".to_string(), l2_pos_emb));
if self.rope.is_some() {
let rope = self
.rope
.as_ref()
.ok_or_else(|| WhisperError::Model("Moonshine decoder requires RoPE".into()))?;
let enc_seq_len = encoder_output.len() / self.d_model;
for (layer_idx, block) in self.moonshine_blocks.iter().enumerate() {
x = block.forward(&x, encoder_output, seq_len, enc_seq_len, rope)?;
let l2: f32 = x.iter().map(|v| v * v).sum::<f32>().sqrt();
let last_start = (seq_len - 1) * self.d_model;
let last_l2: f32 = x[last_start..last_start + self.d_model]
.iter()
.map(|v| v * v)
.sum::<f32>()
.sqrt();
trace.push((format!("layer_{layer_idx}"), l2));
trace.push((format!("layer_{layer_idx}_last"), last_l2));
}
} else {
let causal_mask = MultiHeadAttention::causal_mask(seq_len);
for (layer_idx, block) in self.blocks.iter().enumerate() {
x = block.forward(&x, encoder_output, Some(&causal_mask))?;
let l2: f32 = x.iter().map(|v| v * v).sum::<f32>().sqrt();
let last_start = (seq_len - 1) * self.d_model;
let last_l2: f32 = x[last_start..last_start + self.d_model]
.iter()
.map(|v| v * v)
.sum::<f32>()
.sqrt();
trace.push((format!("layer_{layer_idx}"), l2));
trace.push((format!("layer_{layer_idx}_last"), last_l2));
}
}
let last_start = (seq_len - 1) * self.d_model;
let last_before_ln = &x[last_start..last_start + self.d_model];
let mean: f32 = last_before_ln.iter().sum::<f32>() / self.d_model as f32;
let variance: f32 = last_before_ln
.iter()
.map(|&v| (v - mean).powi(2))
.sum::<f32>()
/ self.d_model as f32;
let std = (variance + 1e-5_f32).sqrt();
trace.push(("ln_mean".to_string(), mean));
trace.push(("ln_var".to_string(), variance));
trace.push(("ln_std".to_string(), std));
if let Some(ref rms) = self.ln_post_rms {
x = rms.forward(&x, seq_len)?;
} else {
let ln_w_l2: f32 = self
.ln_post
.weight
.iter()
.map(|v| v * v)
.sum::<f32>()
.sqrt();
let ln_w_mean: f32 =
self.ln_post.weight.iter().sum::<f32>() / self.ln_post.weight.len() as f32;
let ln_b_l2: f32 = self.ln_post.bias.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("ln_weight_l2".to_string(), ln_w_l2));
trace.push(("ln_weight_mean".to_string(), ln_w_mean));
trace.push(("ln_bias_l2".to_string(), ln_b_l2));
x = self.ln_post.forward(&x)?;
}
let l2_post_ln: f32 = x.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("post_ln".to_string(), l2_post_ln));
let last_hidden_start = (seq_len - 1) * self.d_model;
let last_hidden = &x[last_hidden_start..last_hidden_start + self.d_model];
let l2_last_hidden: f32 = last_hidden.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("last_hidden".to_string(), l2_last_hidden));
let logits = self.project_to_vocab(&x, seq_len);
let l2_logits: f32 = logits.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("logits".to_string(), l2_logits));
let last_logits = &logits[(seq_len - 1) * self.n_vocab..];
let l2_last_logits: f32 = last_logits.iter().map(|v| v * v).sum::<f32>().sqrt();
trace.push(("last_logits".to_string(), l2_last_logits));
Ok((logits, trace))
}
fn embed_tokens(&self, tokens: &[u32]) -> WhisperResult<Vec<f32>> {
let seq_len = tokens.len();
let mut embeddings = vec![0.0_f32; seq_len * self.d_model];
for (pos, &token) in tokens.iter().enumerate() {
let token_idx = token as usize;
if token_idx >= self.n_vocab {
return Err(WhisperError::Model(format!(
"token {} out of vocabulary range {}",
token, self.n_vocab
)));
}
let emb_start = token_idx * self.d_model;
let out_start = pos * self.d_model;
embeddings[out_start..out_start + self.d_model]
.copy_from_slice(&self.token_embedding[emb_start..emb_start + self.d_model]);
}
Ok(embeddings)
}
fn project_to_vocab(&self, x: &[f32], seq_len: usize) -> Vec<f32> {
if seq_len == 1 {
if let Some(ref emb_f16) = self.token_embedding_f16 {
return crate::simd::tiled_matvec_f16(emb_f16, x, self.n_vocab, self.d_model);
}
}
let Ok(x_matrix) = Matrix::from_slice(seq_len, self.d_model, x) else {
return vec![0.0; seq_len * self.n_vocab];
};
x_matrix
.matmul(&self.token_embedding_transposed)
.map_or_else(
|_| vec![0.0; seq_len * self.n_vocab],
|logits| logits.as_slice().to_vec(),
)
}
fn project_to_vocab_into(&self, x: &[f32], out: &mut [f32]) {
debug_assert_eq!(out.len(), self.n_vocab, "logits buffer size mismatch");
if let Some(ref emb_f16) = self.token_embedding_f16 {
crate::simd::tiled_matvec_f16_into(emb_f16, x, out, self.n_vocab, self.d_model);
} else {
let Ok(x_matrix) = Matrix::from_slice(1, self.d_model, x) else {
out.fill(0.0);
return;
};
match x_matrix.matmul(&self.token_embedding_transposed) {
Ok(logits) => {
let src = logits.as_slice();
let copy_len = src.len().min(out.len());
out[..copy_len].copy_from_slice(&src[..copy_len]);
}
Err(_) => out.fill(0.0),
}
}
}
#[must_use]
pub fn project_to_vocab_debug(&self, hidden: &[f32]) -> Vec<f32> {
self.project_to_vocab(hidden, 1)
}
#[must_use]
pub const fn n_layers(&self) -> usize {
self.n_layers
}
#[must_use]
pub const fn d_model(&self) -> usize {
self.d_model
}
#[must_use]
pub const fn n_heads(&self) -> usize {
self.n_heads
}
#[must_use]
pub const fn n_vocab(&self) -> usize {
self.n_vocab
}
#[must_use]
pub const fn max_len(&self) -> usize {
self.max_len
}
#[must_use]
pub fn blocks(&self) -> &[DecoderBlock] {
&self.blocks
}
#[must_use]
pub fn token_embedding(&self) -> &[f32] {
&self.token_embedding
}
pub fn token_embedding_mut(&mut self) -> &mut [f32] {
&mut self.token_embedding
}
#[must_use]
pub fn positional_embedding(&self) -> &[f32] {
&self.positional_embedding
}
pub fn positional_embedding_mut(&mut self) -> &mut [f32] {
&mut self.positional_embedding
}
#[must_use]
pub fn positional_encoding(&self) -> PositionalEncoding {
self.positional_encoding
}
#[must_use]
pub fn attention_type(&self) -> AttentionType {
self.attention_type
}
pub fn blocks_mut(&mut self) -> &mut [DecoderBlock] {
&mut self.blocks
}
#[must_use]
pub fn moonshine_blocks(&self) -> &[MoonshineDecoderBlock] {
&self.moonshine_blocks
}
pub fn moonshine_blocks_mut(&mut self) -> &mut [MoonshineDecoderBlock] {
&mut self.moonshine_blocks
}
#[must_use]
pub fn rope(&self) -> Option<&RotaryEmbedding> {
self.rope.as_ref()
}
#[must_use]
pub fn ln_post_rms(&self) -> Option<&LayerNormNoBias> {
self.ln_post_rms.as_ref()
}
pub fn ln_post_rms_mut(&mut self) -> Option<&mut LayerNormNoBias> {
self.ln_post_rms.as_mut()
}
#[must_use]
pub fn ln_post(&self) -> &crate::model::encoder::LayerNorm {
&self.ln_post
}
pub fn ln_post_mut(&mut self) -> &mut crate::model::encoder::LayerNorm {
&mut self.ln_post
}
#[must_use]
pub fn create_kv_cache(&self) -> DecoderKVCache {
match self.attention_type {
AttentionType::Gqa { kv_heads } => {
let head_dim = self.d_model / self.n_heads;
let padded_hd = if self.moonshine_blocks.is_empty() {
head_dim
} else {
self.moonshine_blocks[0].self_attn.config.padded_head_dim()
};
let kv_dim = kv_heads as usize * padded_hd;
DecoderKVCache::new_gqa(self.n_layers, kv_dim, self.d_model, self.max_len)
}
AttentionType::Mha => DecoderKVCache::new(self.n_layers, self.d_model, self.max_len),
}
}
#[must_use]
pub fn create_decoder_scratch(&self) -> DecoderScratch {
let d_ff = self.d_model * 4;
let d_head = self.d_model / self.n_heads;
let attn_max_kv = self.max_len.max(1500);
DecoderScratch::new_with_attn(self.d_model, d_ff, d_head, attn_max_kv, self.n_vocab)
}
#[must_use]
pub fn create_decoder_scratch_with_enc_len(&self, enc_ctx_len: usize) -> DecoderScratch {
let d_ff = self.d_model * 4;
let d_head = self.d_model / self.n_heads;
let attn_max_kv = self.max_len.max(enc_ctx_len);
DecoderScratch::new_with_attn(self.d_model, d_ff, d_head, attn_max_kv, self.n_vocab)
}
#[cfg(feature = "realizar-inference")]
#[must_use]
pub fn create_paged_kv_cache(&self, total_pages: usize) -> PagedDecoderKVCache {
let config = ModelConfig {
model_type: crate::model::ModelType::Tiny, n_vocab: self.n_vocab as u32,
n_audio_ctx: 1500,
n_audio_state: self.d_model as u32,
n_audio_head: self.n_heads as u32,
n_audio_layer: self.n_layers as u32,
n_text_ctx: self.max_len as u32,
n_text_state: self.d_model as u32,
n_text_head: self.n_heads as u32,
n_text_layer: self.n_layers as u32,
n_mels: 80,
audio_frontend: crate::model::AudioFrontend::MelFilterbank,
positional_encoding: crate::model::PositionalEncoding::Sinusoidal,
ffn_activation: crate::format::FfnActivation::Gelu,
attention_type: crate::model::AttentionType::Mha,
model_family: crate::format::ModelFamily::Whisper,
};
PagedDecoderKVCache::new(&config, total_pages)
}
#[cfg(feature = "realizar-inference")]
pub fn forward_one_paged(
&self,
token: u32,
encoder_output: &[f32],
cache: &mut PagedDecoderKVCache,
seq_id: crate::realizar_inference::SeqId,
) -> WhisperResult<Vec<f32>> {
if self.rope.is_some() {
return Err(WhisperError::Model(
"forward_one_paged is not supported for Moonshine; use forward_one instead".into(),
));
}
let pos = cache.seq_len(seq_id);
if pos >= self.max_len {
return Err(WhisperError::Model(format!(
"cache position {} exceeds max {}",
pos, self.max_len
)));
}
if token as usize >= self.n_vocab {
return Err(WhisperError::Model(format!(
"token {} out of vocabulary range {}",
token, self.n_vocab
)));
}
let emb_start = (token as usize) * self.d_model;
let mut x: Vec<f32> = self.token_embedding[emb_start..emb_start + self.d_model].to_vec();
let pos_start = pos * self.d_model;
for (x_val, pos_emb) in x
.iter_mut()
.zip(&self.positional_embedding[pos_start..pos_start + self.d_model])
{
*x_val += pos_emb;
}
for (layer_idx, block) in self.blocks.iter().enumerate() {
x =
self.forward_block_paged(block, &x, encoder_output, layer_idx, cache, seq_id, pos)?;
}
cache.increment_seq_len(seq_id);
let x = self.ln_post.forward(&x)?;
Ok(self.project_to_vocab(&x, 1))
}
#[cfg(feature = "realizar-inference")]
#[allow(clippy::too_many_arguments)]
fn forward_block_paged(
&self,
block: &DecoderBlock,
x: &[f32],
encoder_output: &[f32],
layer_idx: usize,
cache: &mut PagedDecoderKVCache,
seq_id: crate::realizar_inference::SeqId,
pos: usize,
) -> WhisperResult<Vec<f32>> {
let normed = block.ln1.forward(x)?;
let q = block.self_attn.w_q().forward_simd(&normed, 1)?;
let k_new = block.self_attn.w_k().forward_simd(&normed, 1)?;
let v_new = block.self_attn.w_v().forward_simd(&normed, 1)?;
cache.append(seq_id, layer_idx, &k_new, &v_new)?;
let (k_full, v_full) = cache.get_all_kv_n(seq_id, layer_idx, pos + 1)?;
let attn_out = self.compute_attention_cached(&block.self_attn, &q, &k_full, &v_full)?;
let attn_out = block.self_attn.w_o().forward_simd(&attn_out, 1)?;
let mut residual: Vec<f32> = x.iter().zip(attn_out.iter()).map(|(a, b)| a + b).collect();
let normed = block.ln2.forward(&residual)?;
let enc_len = encoder_output.len() / self.d_model;
let k_enc = block
.cross_attn
.w_k()
.forward_simd(encoder_output, enc_len)?;
let v_enc = block
.cross_attn
.w_v()
.forward_simd(encoder_output, enc_len)?;
let q = block.cross_attn.w_q().forward_simd(&normed, 1)?;
let cross_out = self.compute_attention_cached(&block.cross_attn, &q, &k_enc, &v_enc)?;
let cross_out = block.cross_attn.w_o().forward_simd(&cross_out, 1)?;
for (r, c) in residual.iter_mut().zip(cross_out.iter()) {
*r += c;
}
let normed = block.ln3.forward(&residual)?;
let ffn_out = block.ffn.forward(&normed)?;
for (r, f) in residual.iter_mut().zip(ffn_out.iter()) {
*r += f;
}
Ok(residual)
}
#[cfg(feature = "realizar-inference")]
pub fn generate_paged(
&self,
encoder_output: &[f32],
initial_tokens: &[u32],
max_tokens: usize,
eos_token: u32,
) -> WhisperResult<Vec<u32>> {
let pages_needed = max_tokens.div_ceil(16) + 1; let mut cache = self.create_paged_kv_cache(pages_needed);
let seq_id = cache.allocate_sequence(0)?;
let mut tokens = initial_tokens.to_vec();
for &token in initial_tokens {
let _ = self.forward_one_paged(token, encoder_output, &mut cache, seq_id)?;
}
for _ in initial_tokens.len()..max_tokens {
let last_token = *tokens
.last()
.ok_or_else(|| WhisperError::Model("empty token sequence".into()))?;
let logits = self.forward_one_paged(last_token, encoder_output, &mut cache, seq_id)?;
let next_token = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map_or(eos_token, |(idx, _)| idx as u32);
tokens.push(next_token);
if next_token == eos_token {
break;
}
}
Ok(tokens)
}
#[allow(clippy::needless_range_loop)]
pub fn forward_one(
&self,
token: u32,
encoder_output: &[f32],
cache: &mut DecoderKVCache,
) -> WhisperResult<Vec<f32>> {
let pos = cache.seq_len();
if pos >= self.max_len {
return Err(WhisperError::Model(format!(
"cache position {} exceeds max {}",
pos, self.max_len
)));
}
if token as usize >= self.n_vocab {
return Err(WhisperError::Model(format!(
"token {} out of vocabulary range {}",
token, self.n_vocab
)));
}
let emb_start = (token as usize) * self.d_model;
let mut x: Vec<f32> = self.token_embedding[emb_start..emb_start + self.d_model].to_vec();
if self.rope.is_some() {
let rope = self
.rope
.as_ref()
.ok_or_else(|| WhisperError::Model("Moonshine decoder requires RoPE".into()))?;
let pos = cache.seq_position;
let enc_seq_len = encoder_output.len() / self.d_model;
for (layer_idx, block) in self.moonshine_blocks.iter().enumerate() {
x = block.forward_cached(
&x,
encoder_output,
enc_seq_len,
pos,
rope,
&mut cache.self_attn_cache[layer_idx],
&mut cache.cross_attn_cache[layer_idx],
cache.cross_attn_cached,
)?;
}
if !cache.cross_attn_cached {
cache.cross_attn_cached = true;
}
cache.seq_position += 1;
if let Some(ref rms) = self.ln_post_rms {
x = rms.forward(&x, 1)?;
}
Ok(self.project_to_vocab(&x, 1))
} else {
let pos_start = pos * self.d_model;
for d in 0..self.d_model {
x[d] += self.positional_embedding[pos_start + d];
}
for (layer_idx, block) in self.blocks.iter().enumerate() {
x = self.forward_block_cached(block, &x, encoder_output, layer_idx, cache)?;
}
if !cache.cross_attn_cached {
cache.cross_attn_cached = true;
}
let x = self.ln_post.forward(&x)?;
Ok(self.project_to_vocab(&x, 1))
}
}
#[allow(clippy::needless_range_loop)]
pub fn forward_one_with_scratch(
&self,
token: u32,
encoder_output: &[f32],
cache: &mut DecoderKVCache,
scratch: &mut DecoderScratch,
) -> WhisperResult<Vec<f32>> {
if self.rope.is_some() {
return self.forward_one(token, encoder_output, cache);
}
let pos = cache.seq_len();
if pos >= self.max_len {
return Err(WhisperError::Model(format!(
"cache position {} exceeds max {}",
pos, self.max_len
)));
}
if token as usize >= self.n_vocab {
return Err(WhisperError::Model(format!(
"token {} out of vocabulary range {}",
token, self.n_vocab
)));
}
let emb_start = (token as usize) * self.d_model;
let mut x: Vec<f32> = self.token_embedding[emb_start..emb_start + self.d_model].to_vec();
let pos_start = pos * self.d_model;
for d in 0..self.d_model {
x[d] += self.positional_embedding[pos_start + d];
}
for (layer_idx, block) in self.blocks.iter().enumerate() {
self.forward_block_cached_with_scratch(
block,
&mut x,
encoder_output,
layer_idx,
cache,
scratch,
)?;
}
if !cache.cross_attn_cached {
cache.cross_attn_cached = true;
}
self.ln_post.forward_into(&x, &mut scratch.ln_post_out)?;
self.project_to_vocab_into(&scratch.ln_post_out, &mut scratch.logits);
Ok(scratch.logits.clone())
}
fn forward_block_cached_with_scratch(
&self,
block: &DecoderBlock,
x: &mut [f32],
encoder_output: &[f32],
layer_idx: usize,
cache: &mut DecoderKVCache,
scratch: &mut DecoderScratch,
) -> WhisperResult<()> {
block.ln1.forward_into(x, &mut scratch.normed)?;
let d = self.d_model;
block
.self_attn
.forward_qkv_into(&scratch.normed, &mut scratch.qkv)?;
scratch.q.copy_from_slice(&scratch.qkv[..d]);
scratch.k_new.copy_from_slice(&scratch.qkv[d..2 * d]);
scratch.v_new.copy_from_slice(&scratch.qkv[2 * d..3 * d]);
cache.self_attn_cache[layer_idx].append(&scratch.k_new, &scratch.v_new)?;
let k_full = cache.self_attn_cache[layer_idx].get_key();
let v_full = cache.self_attn_cache[layer_idx].get_value();
self.compute_attention_cached_with_scratch(
&block.self_attn,
&scratch.q,
k_full,
v_full,
&mut scratch.attn,
)?;
block
.self_attn
.w_o()
.forward_simd_into(&scratch.attn.output, 1, &mut scratch.sa_proj)?;
for (xi, &pi) in x.iter_mut().zip(scratch.sa_proj.iter()) {
*xi += pi;
}
block.ln2.forward_into(x, &mut scratch.normed2)?;
if !cache.cross_attn_cached || cache.cross_attn_cache[layer_idx].is_empty() {
let enc_len = encoder_output.len() / self.d_model;
let k_enc = block
.cross_attn
.w_k()
.forward_simd(encoder_output, enc_len)?;
let v_enc = block
.cross_attn
.w_v()
.forward_simd(encoder_output, enc_len)?;
cache.cross_attn_cache[layer_idx].append(&k_enc, &v_enc)?;
block
.cross_attn
.w_q()
.forward_simd_into(&scratch.normed2, 1, &mut scratch.cross_q)?;
self.compute_attention_cached_with_scratch(
&block.cross_attn,
&scratch.cross_q,
&k_enc,
&v_enc,
&mut scratch.attn,
)?;
block.cross_attn.w_o().forward_simd_into(
&scratch.attn.output,
1,
&mut scratch.cross_proj,
)?;
} else {
let k_cached = cache.cross_attn_cache[layer_idx].get_key();
let v_cached = cache.cross_attn_cache[layer_idx].get_value();
block
.cross_attn
.w_q()
.forward_simd_into(&scratch.normed2, 1, &mut scratch.cross_q)?;
self.compute_attention_cached_with_scratch(
&block.cross_attn,
&scratch.cross_q,
k_cached,
v_cached,
&mut scratch.attn,
)?;
block.cross_attn.w_o().forward_simd_into(
&scratch.attn.output,
1,
&mut scratch.cross_proj,
)?;
}
for (xi, &ci) in x.iter_mut().zip(scratch.cross_proj.iter()) {
*xi += ci;
}
block.ln3.forward_into(x, &mut scratch.normed3)?;
block.ffn.forward_into(
&scratch.normed3,
&mut scratch.ffn_hidden,
&mut scratch.ffn_out,
)?;
for (xi, &fi) in x.iter_mut().zip(scratch.ffn_out.iter()) {
*xi += fi;
}
Ok(())
}
pub fn forward_one_hidden(
&self,
token: u32,
encoder_output: &[f32],
cache: &mut DecoderKVCache,
) -> WhisperResult<Vec<f32>> {
let pos = cache.seq_len();
if pos >= self.max_len {
return Err(WhisperError::Model(format!(
"cache position {} exceeds max {}",
pos, self.max_len
)));
}
if token as usize >= self.n_vocab {
return Err(WhisperError::Model(format!(
"token {} out of vocabulary range {}",
token, self.n_vocab
)));
}
let emb_start = (token as usize) * self.d_model;
let mut x: Vec<f32> = self.token_embedding[emb_start..emb_start + self.d_model].to_vec();
if self.rope.is_some() {
let rope = self
.rope
.as_ref()
.ok_or_else(|| WhisperError::Model("Moonshine decoder requires RoPE".into()))?;
let moon_pos = cache.seq_position;
let enc_seq_len = encoder_output.len() / self.d_model;
for (layer_idx, block) in self.moonshine_blocks.iter().enumerate() {
x = block.forward_cached(
&x,
encoder_output,
enc_seq_len,
moon_pos,
rope,
&mut cache.self_attn_cache[layer_idx],
&mut cache.cross_attn_cache[layer_idx],
cache.cross_attn_cached,
)?;
}
if !cache.cross_attn_cached {
cache.cross_attn_cached = true;
}
cache.seq_position += 1;
if let Some(ref rms) = self.ln_post_rms {
x = rms.forward(&x, 1)?;
}
Ok(x)
} else {
let pos_start = pos * self.d_model;
for (x_elem, pos_emb) in x
.iter_mut()
.zip(&self.positional_embedding[pos_start..pos_start + self.d_model])
{
*x_elem += pos_emb;
}
for (layer_idx, block) in self.blocks.iter().enumerate() {
x = self.forward_block_cached(block, &x, encoder_output, layer_idx, cache)?;
}
if !cache.cross_attn_cached {
cache.cross_attn_cached = true;
}
self.ln_post.forward(&x)
}
}
#[cfg(feature = "realizar-inference")]
#[allow(clippy::needless_range_loop)]
pub fn forward_one_fused(
&self,
token: u32,
encoder_output: &[f32],
cache: &mut DecoderKVCache,
) -> WhisperResult<Vec<f32>> {
if self.rope.is_some() {
return Err(WhisperError::Model(
"forward_one_fused is not supported for Moonshine; use forward_one instead".into(),
));
}
let pos = cache.seq_len();
if pos >= self.max_len {
return Err(WhisperError::Model(format!(
"cache position {pos} exceeds max {}",
self.max_len
)));
}
if token as usize >= self.n_vocab {
return Err(WhisperError::Model(format!(
"token {token} out of vocabulary range {}",
self.n_vocab
)));
}
let emb_start = (token as usize) * self.d_model;
let mut x: Vec<f32> = self.token_embedding[emb_start..emb_start + self.d_model].to_vec();
let pos_start = pos * self.d_model;
for d in 0..self.d_model {
x[d] += self.positional_embedding[pos_start + d];
}
for (layer_idx, block) in self.blocks.iter().enumerate() {
x = self.forward_block_cached_fused(block, &x, encoder_output, layer_idx, cache)?;
}
if !cache.cross_attn_cached {
cache.cross_attn_cached = true;
}
let x = self.ln_post.forward(&x)?;
Ok(self.project_to_vocab(&x, 1))
}
fn forward_block_cached(
&self,
block: &DecoderBlock,
x: &[f32],
encoder_output: &[f32],
layer_idx: usize,
cache: &mut DecoderKVCache,
) -> WhisperResult<Vec<f32>> {
let normed = block.ln1.forward(x)?;
let q = block.self_attn.w_q().forward_simd(&normed, 1)?;
let k_new = block.self_attn.w_k().forward_simd(&normed, 1)?;
let v_new = block.self_attn.w_v().forward_simd(&normed, 1)?;
cache.self_attn_cache[layer_idx].append(&k_new, &v_new)?;
let k_full = cache.self_attn_cache[layer_idx].get_key();
let v_full = cache.self_attn_cache[layer_idx].get_value();
let attn_out = self.compute_attention_cached(&block.self_attn, &q, k_full, v_full)?;
let attn_out = block.self_attn.w_o().forward_simd(&attn_out, 1)?;
let mut residual: Vec<f32> = x.iter().zip(attn_out.iter()).map(|(a, b)| a + b).collect();
let normed = block.ln2.forward(&residual)?;
let cross_out = if !cache.cross_attn_cached || cache.cross_attn_cache[layer_idx].is_empty()
{
let enc_len = encoder_output.len() / self.d_model;
let k_enc = block
.cross_attn
.w_k()
.forward_simd(encoder_output, enc_len)?;
let v_enc = block
.cross_attn
.w_v()
.forward_simd(encoder_output, enc_len)?;
cache.cross_attn_cache[layer_idx].append(&k_enc, &v_enc)?;
let q = block.cross_attn.w_q().forward_simd(&normed, 1)?;
let attn_out = self.compute_attention_cached(&block.cross_attn, &q, &k_enc, &v_enc)?;
block.cross_attn.w_o().forward_simd(&attn_out, 1)?
} else {
let k_cached = cache.cross_attn_cache[layer_idx].get_key();
let v_cached = cache.cross_attn_cache[layer_idx].get_value();
let q = block.cross_attn.w_q().forward_simd(&normed, 1)?;
let attn_out =
self.compute_attention_cached(&block.cross_attn, &q, k_cached, v_cached)?;
block.cross_attn.w_o().forward_simd(&attn_out, 1)?
};
for (r, c) in residual.iter_mut().zip(cross_out.iter()) {
*r += c;
}
let normed = block.ln3.forward(&residual)?;
let ffn_out = block.ffn.forward(&normed)?;
for (r, f) in residual.iter_mut().zip(ffn_out.iter()) {
*r += f;
}
Ok(residual)
}
#[cfg(feature = "realizar-inference")]
fn forward_block_cached_fused(
&self,
block: &DecoderBlock,
x: &[f32],
encoder_output: &[f32],
layer_idx: usize,
cache: &mut DecoderKVCache,
) -> WhisperResult<Vec<f32>> {
let normed = block.ln1.forward(x)?;
let q = block.self_attn.w_q().forward_simd(&normed, 1)?;
let k_new = block.self_attn.w_k().forward_simd(&normed, 1)?;
let v_new = block.self_attn.w_v().forward_simd(&normed, 1)?;
cache.self_attn_cache[layer_idx].append(&k_new, &v_new)?;
let k_full = cache.self_attn_cache[layer_idx].get_key();
let v_full = cache.self_attn_cache[layer_idx].get_value();
let attn_out = self.compute_attention_cached(&block.self_attn, &q, k_full, v_full)?;
let attn_out = block.self_attn.w_o().forward_simd(&attn_out, 1)?;
let mut residual: Vec<f32> = x.iter().zip(attn_out.iter()).map(|(a, b)| a + b).collect();
let normed = block.ln2.forward(&residual)?;
let cross_out = if !cache.cross_attn_cached || cache.cross_attn_cache[layer_idx].is_empty()
{
let enc_len = encoder_output.len() / self.d_model;
let k_enc = block
.cross_attn
.w_k()
.forward_simd(encoder_output, enc_len)?;
let v_enc = block
.cross_attn
.w_v()
.forward_simd(encoder_output, enc_len)?;
cache.cross_attn_cache[layer_idx].append(&k_enc, &v_enc)?;
let q = block.cross_attn.w_q().forward_simd(&normed, 1)?;
let attn_out = self.compute_attention_cached(&block.cross_attn, &q, &k_enc, &v_enc)?;
block.cross_attn.w_o().forward_simd(&attn_out, 1)?
} else {
let k_cached = cache.cross_attn_cache[layer_idx].get_key();
let v_cached = cache.cross_attn_cache[layer_idx].get_value();
let q = block.cross_attn.w_q().forward_simd(&normed, 1)?;
let attn_out =
self.compute_attention_cached(&block.cross_attn, &q, k_cached, v_cached)?;
block.cross_attn.w_o().forward_simd(&attn_out, 1)?
};
for (r, c) in residual.iter_mut().zip(cross_out.iter()) {
*r += c;
}
let fused = block.create_fused_ffn()?;
let ffn_out = fused.forward(&residual)?;
for (r, f) in residual.iter_mut().zip(ffn_out.iter()) {
*r += f;
}
Ok(residual)
}
fn compute_attention_cached(
&self,
attn: &MultiHeadAttention,
q: &[f32],
k: &[f32],
v: &[f32],
) -> WhisperResult<Vec<f32>> {
let n_heads = attn.n_heads();
let d_head = attn.d_head();
let kv_len = k.len() / self.d_model;
let mut output = vec![0.0_f32; self.d_model];
let mut q_head = vec![0.0_f32; d_head];
let mut k_head = vec![0.0_f32; kv_len * d_head];
let mut v_head = vec![0.0_f32; kv_len * d_head];
for head in 0..n_heads {
for d in 0..d_head {
q_head[d] = q[head * d_head + d];
}
for pos in 0..kv_len {
for d in 0..d_head {
k_head[pos * d_head + d] = k[pos * self.d_model + head * d_head + d];
v_head[pos * d_head + d] = v[pos * self.d_model + head * d_head + d];
}
}
let head_out =
attn.scaled_dot_product_attention_simd(&q_head, &k_head, &v_head, None)?;
for d in 0..d_head {
output[head * d_head + d] = head_out[d];
}
}
Ok(output)
}
fn compute_attention_cached_with_scratch(
&self,
mha: &MultiHeadAttention,
q: &[f32],
k: &[f32],
v: &[f32],
attn_scratch: &mut AttentionScratch,
) -> WhisperResult<()> {
let n_heads = mha.n_heads();
let d_head = mha.d_head();
let kv_len = k.len() / self.d_model;
let scale = 1.0 / (d_head as f32).sqrt();
for head in 0..n_heads {
let head_offset = head * d_head;
for pos in 0..kv_len {
let src_base = pos * self.d_model + head_offset;
let dst_base = pos * d_head;
attn_scratch.k_head[dst_base..dst_base + d_head]
.copy_from_slice(&k[src_base..src_base + d_head]);
attn_scratch.v_head[dst_base..dst_base + d_head]
.copy_from_slice(&v[src_base..src_base + d_head]);
}
let q_head = &q[head_offset..head_offset + d_head];
for pos in 0..kv_len {
let k_start = pos * d_head;
let mut dot = 0.0_f32;
for d in 0..d_head {
dot += q_head[d] * attn_scratch.k_head[k_start + d];
}
attn_scratch.scores[pos] = dot * scale;
}
crate::simd::softmax_online_inplace(
&attn_scratch.scores[..kv_len],
&mut attn_scratch.weights[..kv_len],
);
let out_start = head * d_head;
for d in 0..d_head {
attn_scratch.output[out_start + d] = 0.0;
}
for pos in 0..kv_len {
let weight = attn_scratch.weights[pos];
let v_start = pos * d_head;
for d in 0..d_head {
attn_scratch.output[out_start + d] += weight * attn_scratch.v_head[v_start + d];
}
}
}
Ok(())
}
pub fn generate(
&self,
encoder_output: &[f32],
initial_tokens: &[u32],
max_tokens: usize,
eos_token: u32,
) -> WhisperResult<Vec<u32>> {
let mut cache = self.create_kv_cache();
let mut scratch = self.create_decoder_scratch();
let mut tokens = initial_tokens.to_vec();
for &token in initial_tokens {
let _ =
self.forward_one_with_scratch(token, encoder_output, &mut cache, &mut scratch)?;
}
for _ in initial_tokens.len()..max_tokens {
let last_token = *tokens
.last()
.ok_or_else(|| WhisperError::Model("empty token sequence".into()))?;
let logits = self.forward_one_with_scratch(
last_token,
encoder_output,
&mut cache,
&mut scratch,
)?;
let next_token = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map_or(eos_token, |(idx, _)| idx as u32);
tokens.push(next_token);
if next_token == eos_token {
break;
}
}
Ok(tokens)
}
#[must_use]
pub fn create_batch_cache(&self, batch_size: usize) -> BatchDecoderCache {
BatchDecoderCache::new(batch_size, self.n_layers, self.d_model, self.max_len)
}
pub fn forward_batch(
&self,
tokens_batch: &[Vec<u32>],
encoder_outputs: &[Vec<f32>],
) -> WhisperResult<BatchDecoderOutput> {
if tokens_batch.is_empty() {
return Err(WhisperError::Model("empty batch".into()));
}
if tokens_batch.len() != encoder_outputs.len() {
return Err(WhisperError::Model(format!(
"batch size mismatch: {} tokens vs {} encoders",
tokens_batch.len(),
encoder_outputs.len()
)));
}
let mut logits = Vec::with_capacity(tokens_batch.len());
let mut seq_lengths = Vec::with_capacity(tokens_batch.len());
for (tokens, encoder_out) in tokens_batch.iter().zip(encoder_outputs.iter()) {
let item_logits = self.forward(tokens, encoder_out)?;
seq_lengths.push(tokens.len());
logits.push(item_logits);
}
Ok(BatchDecoderOutput {
logits,
seq_lengths,
})
}
pub fn forward_one_batch(
&self,
tokens: &[u32],
encoder_outputs: &[Vec<f32>],
cache: &mut BatchDecoderCache,
) -> WhisperResult<BatchDecoderOutput> {
let batch_size = cache.batch_size();
if tokens.len() != batch_size {
return Err(WhisperError::Model(format!(
"token count {} doesn't match batch size {}",
tokens.len(),
batch_size
)));
}
if encoder_outputs.len() != batch_size {
return Err(WhisperError::Model(format!(
"encoder count {} doesn't match batch size {}",
encoder_outputs.len(),
batch_size
)));
}
let mut logits = Vec::with_capacity(batch_size);
for (idx, (&token, encoder_out)) in tokens.iter().zip(encoder_outputs.iter()).enumerate() {
let item_cache = cache
.get_cache_mut(idx)
.ok_or_else(|| WhisperError::Model(format!("cache index {idx} out of bounds")))?;
let item_logits = self.forward_one(token, encoder_out, item_cache)?;
logits.push(item_logits);
}
Ok(BatchDecoderOutput {
logits,
seq_lengths: vec![1; batch_size],
})
}
pub fn forward_one_batch_fused(
&self,
tokens: &[u32],
encoder_outputs: &[Vec<f32>],
cache: &mut BatchDecoderCache,
) -> WhisperResult<BatchDecoderOutput> {
let batch_size = cache.batch_size();
if tokens.len() != batch_size {
return Err(WhisperError::Model(format!(
"token count doesn't match batch size"
)));
}
if encoder_outputs.len() != batch_size {
return Err(WhisperError::Model(format!(
"encoder count doesn't match batch size"
)));
}
let mut x = vec![0.0_f32; batch_size * self.d_model];
for (i, &token) in tokens.iter().enumerate() {
if token as usize >= self.n_vocab {
return Err(WhisperError::Model(format!("token {} out of range", token)));
}
let emb_start = (token as usize) * self.d_model;
let out_start = i * self.d_model;
x[out_start..out_start + self.d_model]
.copy_from_slice(&self.token_embedding[emb_start..emb_start + self.d_model]);
let item_cache = cache
.get_cache(i)
.ok_or_else(|| WhisperError::Model("cache must exist for beam".to_string()))?;
let pos = item_cache.seq_len();
if pos >= self.max_len {
return Err(WhisperError::Model(format!("cache position exceeds max")));
}
let pos_start = pos * self.d_model;
for d in 0..self.d_model {
x[out_start + d] += self.positional_embedding[pos_start + d];
}
}
let mut qkv_batch = vec![0.0_f32; batch_size * 3 * self.d_model];
let mut attn_out_batch = vec![0.0_f32; batch_size * self.d_model];
let mut q_batch = vec![0.0_f32; batch_size * self.d_model];
for (layer_idx, block) in self.blocks.iter().enumerate() {
let mut normed = vec![0.0_f32; batch_size * self.d_model];
for i in 0..batch_size {
let start = i * self.d_model;
block.ln1.forward_into(
&x[start..start + self.d_model],
&mut normed[start..start + self.d_model],
)?;
}
block
.self_attn
.forward_qkv_batch_into(&normed, batch_size, &mut qkv_batch)?;
for i in 0..batch_size {
let start = i * 3 * self.d_model;
let q_i = &qkv_batch[start..start + self.d_model];
let k_i = &qkv_batch[start + self.d_model..start + 2 * self.d_model];
let v_i = &qkv_batch[start + 2 * self.d_model..start + 3 * self.d_model];
let item_cache = cache
.get_cache_mut(i)
.ok_or_else(|| WhisperError::Model("cache must exist for beam".to_string()))?;
item_cache.self_attn_cache[layer_idx].append(k_i, v_i)?;
let k_full = item_cache.self_attn_cache[layer_idx].get_key();
let v_full = item_cache.self_attn_cache[layer_idx].get_value();
let attn_out_i =
self.compute_attention_cached(&block.self_attn, q_i, k_full, v_full)?;
let out_start = i * self.d_model;
attn_out_batch[out_start..out_start + self.d_model].copy_from_slice(&attn_out_i);
}
let mut attn_proj = vec![0.0_f32; batch_size * self.d_model];
block
.self_attn
.w_o()
.forward_simd_into(&attn_out_batch, batch_size, &mut attn_proj)?;
for (a, b) in x.iter_mut().zip(attn_proj.iter()) {
*a += b;
}
for i in 0..batch_size {
let start = i * self.d_model;
block.ln2.forward_into(
&x[start..start + self.d_model],
&mut normed[start..start + self.d_model],
)?;
}
block
.cross_attn
.w_q()
.forward_simd_into(&normed, batch_size, &mut q_batch)?;
for i in 0..batch_size {
let item_cache = cache
.get_cache_mut(i)
.ok_or_else(|| WhisperError::Model("cache must exist for beam".to_string()))?;
let q_i = &q_batch[i * self.d_model..(i + 1) * self.d_model];
let attn_out_i = if !item_cache.cross_attn_cached
|| item_cache.cross_attn_cache[layer_idx].is_empty()
{
let enc_len = encoder_outputs[i].len() / self.d_model;
let k_enc = block
.cross_attn
.w_k()
.forward_simd(&encoder_outputs[i], enc_len)?;
let v_enc = block
.cross_attn
.w_v()
.forward_simd(&encoder_outputs[i], enc_len)?;
item_cache.cross_attn_cache[layer_idx].append(&k_enc, &v_enc)?;
self.compute_attention_cached(&block.cross_attn, q_i, &k_enc, &v_enc)?
} else {
let k_cached = item_cache.cross_attn_cache[layer_idx].get_key();
let v_cached = item_cache.cross_attn_cache[layer_idx].get_value();
self.compute_attention_cached(&block.cross_attn, q_i, k_cached, v_cached)?
};
let out_start = i * self.d_model;
attn_out_batch[out_start..out_start + self.d_model].copy_from_slice(&attn_out_i);
}
block.cross_attn.w_o().forward_simd_into(
&attn_out_batch,
batch_size,
&mut attn_proj,
)?;
for (a, b) in x.iter_mut().zip(attn_proj.iter()) {
*a += b;
}
for i in 0..batch_size {
let start = i * self.d_model;
block.ln3.forward_into(
&x[start..start + self.d_model],
&mut normed[start..start + self.d_model],
)?;
}
let mut ffn_hidden = vec![0.0_f32; batch_size * block.ffn.d_ff];
let mut ffn_out = vec![0.0_f32; batch_size * self.d_model];
block
.ffn
.forward_into(&normed, &mut ffn_hidden, &mut ffn_out)?;
for (a, b) in x.iter_mut().zip(ffn_out.iter()) {
*a += b;
}
}
for i in 0..batch_size {
let item_cache = cache
.get_cache_mut(i)
.ok_or_else(|| WhisperError::Model("cache must exist for beam".to_string()))?;
if !item_cache.cross_attn_cached {
item_cache.cross_attn_cached = true;
}
}
let mut final_normed = vec![0.0_f32; batch_size * self.d_model];
for i in 0..batch_size {
let start = i * self.d_model;
self.ln_post.forward_into(
&x[start..start + self.d_model],
&mut final_normed[start..start + self.d_model],
)?;
}
let logits_flat = self.project_to_vocab(&final_normed, batch_size);
let mut logits = Vec::with_capacity(batch_size);
for i in 0..batch_size {
let start = i * self.n_vocab;
logits.push(logits_flat[start..start + self.n_vocab].to_vec());
}
Ok(BatchDecoderOutput {
logits,
seq_lengths: vec![1; batch_size],
})
}
pub fn generate_batch(
&self,
encoder_outputs: &[Vec<f32>],
initial_tokens: &[Vec<u32>],
max_tokens: usize,
eos_token: u32,
) -> WhisperResult<Vec<Vec<u32>>> {
let batch_size = encoder_outputs.len();
if initial_tokens.len() != batch_size {
return Err(WhisperError::Model(format!(
"initial tokens count {} doesn't match batch size {}",
initial_tokens.len(),
batch_size
)));
}
let mut cache = self.create_batch_cache(batch_size);
let mut sequences: Vec<Vec<u32>> = initial_tokens.to_vec();
let mut finished = vec![false; batch_size];
for (idx, tokens) in initial_tokens.iter().enumerate() {
let item_cache = cache
.get_cache_mut(idx)
.ok_or_else(|| WhisperError::Model(format!("cache index {idx} out of bounds")))?;
for &token in tokens {
let _ = self.forward_one(token, &encoder_outputs[idx], item_cache)?;
}
}
for _ in 0..max_tokens {
if finished.iter().all(|&f| f) {
break;
}
let last_tokens: Vec<u32> = sequences
.iter()
.map(|seq| *seq.last().unwrap_or(&0))
.collect();
let outputs = self.forward_one_batch(&last_tokens, encoder_outputs, &mut cache)?;
for (idx, logits) in outputs.logits.iter().enumerate() {
if finished[idx] {
continue;
}
let next_token = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map_or(eos_token, |(i, _)| i as u32);
sequences[idx].push(next_token);
if next_token == eos_token {
finished[idx] = true;
}
}
}
Ok(sequences)
}
}
#[cfg(feature = "realizar-inference")]
#[allow(dead_code)] pub struct SpeculativeDecoderWrapper<'a> {
decoder: &'a Decoder,
encoder_output: &'a [f32],
cache: std::cell::RefCell<&'a mut DecoderKVCache>,
eos_token_id: u32,
}
#[cfg(feature = "realizar-inference")]
#[allow(dead_code)] impl<'a> SpeculativeDecoderWrapper<'a> {
pub fn new(
decoder: &'a Decoder,
encoder_output: &'a [f32],
cache: &'a mut DecoderKVCache,
) -> Self {
Self {
decoder,
encoder_output,
cache: std::cell::RefCell::new(cache),
eos_token_id: 50257, }
}
pub fn with_eos_token(mut self, eos_token: u32) -> Self {
self.eos_token_id = eos_token;
self
}
}
#[cfg(feature = "realizar-inference")]
impl crate::realizar_inference::SpeculativeModel for SpeculativeDecoderWrapper<'_> {
fn forward(
&self,
tokens: &[u32],
) -> Result<Vec<f32>, crate::realizar_inference::SpeculativeError> {
let mut cache = self.cache.borrow_mut();
if cache.seq_len() > tokens.len() {
cache.clear();
}
for (i, &token) in tokens.iter().enumerate() {
if i >= cache.seq_len() {
self.decoder
.forward_one(token, self.encoder_output, &mut cache)
.map_err(|e| {
crate::realizar_inference::SpeculativeError::TargetModelError(format!(
"decoder forward failed: {e}"
))
})?;
}
}
let last_token = *tokens.last().ok_or_else(|| {
crate::realizar_inference::SpeculativeError::TargetModelError(
"empty token sequence".into(),
)
})?;
let logits = self
.decoder
.forward_one(last_token, self.encoder_output, &mut cache)
.map_err(|e| {
crate::realizar_inference::SpeculativeError::TargetModelError(format!(
"decoder forward failed: {e}"
))
})?;
Ok(logits)
}
fn sample(
&self,
logits: &[f32],
) -> Result<crate::realizar_inference::TokenProb, crate::realizar_inference::SpeculativeError>
{
let (token, &logit) = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.ok_or_else(|| {
crate::realizar_inference::SpeculativeError::TargetModelError("empty logits".into())
})?;
let max_logit = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let log_sum_exp = max_logit
+ logits
.iter()
.map(|&l| (l - max_logit).exp())
.sum::<f32>()
.ln();
let log_prob = logit - log_sum_exp;
Ok(crate::realizar_inference::TokenProb {
token: token as u32,
log_prob,
})
}
fn vocab_size(&self) -> usize {
self.decoder.n_vocab
}
fn eos_token(&self) -> u32 {
self.eos_token_id
}
}
#[cfg(feature = "realizar-inference")]
#[derive(Debug, Clone)]
pub struct WhisperSpeculativeConfig {
pub lookahead: usize,
pub acceptance_threshold: f32,
pub max_tokens: usize,
pub eos_token: u32,
}
#[cfg(feature = "realizar-inference")]
impl Default for WhisperSpeculativeConfig {
fn default() -> Self {
Self {
lookahead: 4, acceptance_threshold: 0.8,
max_tokens: 448, eos_token: 50257,
}
}
}
#[cfg(feature = "realizar-inference")]
impl Decoder {
pub fn generate_speculative(
&self,
draft_decoder: &Decoder,
encoder_output: &[f32],
initial_tokens: &[u32],
config: &WhisperSpeculativeConfig,
) -> WhisperResult<Vec<u32>> {
let mut draft_cache = draft_decoder.create_kv_cache();
let mut target_cache = self.create_kv_cache();
let mut draft_scratch = draft_decoder.create_decoder_scratch();
let mut target_scratch = self.create_decoder_scratch();
let mut tokens = initial_tokens.to_vec();
for &token in initial_tokens {
let _ = draft_decoder.forward_one_with_scratch(
token,
encoder_output,
&mut draft_cache,
&mut draft_scratch,
)?;
let _ = self.forward_one_with_scratch(
token,
encoder_output,
&mut target_cache,
&mut target_scratch,
)?;
}
while tokens.len() < config.max_tokens {
let mut draft_tokens = Vec::with_capacity(config.lookahead);
let mut draft_probs = Vec::with_capacity(config.lookahead);
let mut current_token = *tokens
.last()
.ok_or_else(|| WhisperError::Model("empty token sequence".into()))?;
for _ in 0..config.lookahead {
let draft_logits = draft_decoder.forward_one_with_scratch(
current_token,
encoder_output,
&mut draft_cache,
&mut draft_scratch,
)?;
let (next_token, prob) = sample_with_prob(&draft_logits);
if next_token == config.eos_token {
draft_tokens.push(next_token);
draft_probs.push(prob);
break;
}
draft_tokens.push(next_token);
draft_probs.push(prob);
current_token = next_token;
}
if draft_tokens.is_empty() {
break;
}
let mut accepted_count = 0;
for (i, &draft_token) in draft_tokens.iter().enumerate() {
let prev_token = if i == 0 {
*tokens
.last()
.ok_or_else(|| WhisperError::Model("empty token sequence".into()))?
} else {
draft_tokens[i - 1]
};
let target_logits = self.forward_one_with_scratch(
prev_token,
encoder_output,
&mut target_cache,
&mut target_scratch,
)?;
let (target_token, target_prob) = sample_with_prob(&target_logits);
if draft_token == target_token {
tokens.push(draft_token);
accepted_count += 1;
if draft_token == config.eos_token {
return Ok(tokens);
}
} else {
tokens.push(target_token);
draft_cache.clear();
for &t in &tokens[..tokens.len() - 1] {
let _ = draft_decoder.forward_one_with_scratch(
t,
encoder_output,
&mut draft_cache,
&mut draft_scratch,
)?;
}
if target_token == config.eos_token {
return Ok(tokens);
}
break;
}
let _ = target_prob;
let _ = draft_probs[i];
}
if accepted_count < draft_tokens.len() {
}
}
Ok(tokens)
}
}
#[cfg(feature = "realizar-inference")]
fn sample_with_prob(logits: &[f32]) -> (u32, f32) {
let (token, &logit) = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or((0, &0.0));
let max_logit = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let exp_sum: f32 = logits.iter().map(|&l| (l - max_logit).exp()).sum();
let prob = if exp_sum > 0.0 {
(logit - max_logit).exp() / exp_sum
} else {
0.0
};
(token as u32, prob)
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
fn create_decoder_with_test_weights(config: &ModelConfig) -> Decoder {
let mut decoder = Decoder::new(config);
let d_model = config.n_text_state as usize;
for block in decoder.blocks_mut() {
let scale = 0.1_f32;
let q_weight: Vec<f32> = (0..d_model * d_model)
.map(|i| {
let row = i / d_model;
let col = i % d_model;
if row == col {
scale
} else {
scale * ((i as f32 * 0.01).sin() * 0.1)
}
})
.collect();
block.cross_attn_mut().w_q_mut().set_weight(&q_weight);
let k_weight: Vec<f32> = (0..d_model * d_model)
.map(|i| {
let row = i / d_model;
let col = i % d_model;
if row == col {
scale
} else {
scale * ((i as f32 * 0.02).cos() * 0.1)
}
})
.collect();
block.cross_attn_mut().w_k_mut().set_weight(&k_weight);
let v_weight: Vec<f32> = (0..d_model * d_model)
.map(|i| {
let row = i / d_model;
let col = i % d_model;
if row == col {
scale
} else {
scale * ((i as f32 * 0.03).sin() * 0.1)
}
})
.collect();
block.cross_attn_mut().w_v_mut().set_weight(&v_weight);
let o_weight: Vec<f32> = (0..d_model * d_model)
.map(|i| {
let row = i / d_model;
let col = i % d_model;
if row == col {
scale
} else {
scale * ((i as f32 * 0.04).cos() * 0.1)
}
})
.collect();
block.cross_attn_mut().w_o_mut().set_weight(&o_weight);
let self_attn_weight: Vec<f32> = (0..d_model * d_model)
.map(|i| {
let row = i / d_model;
let col = i % d_model;
if row == col {
scale
} else {
0.0
}
})
.collect();
block
.self_attn_mut()
.w_q_mut()
.set_weight(&self_attn_weight);
block
.self_attn_mut()
.w_k_mut()
.set_weight(&self_attn_weight);
block
.self_attn_mut()
.w_v_mut()
.set_weight(&self_attn_weight);
block
.self_attn_mut()
.w_o_mut()
.set_weight(&self_attn_weight);
let d_ff = d_model * 4;
let fc1_weight: Vec<f32> = (0..d_ff * d_model)
.map(|i| (i as f32 * 0.001).sin() * 0.1)
.collect();
block.ffn.fc1.set_weight(&fc1_weight);
let fc2_weight: Vec<f32> = (0..d_model * d_ff)
.map(|i| (i as f32 * 0.002).cos() * 0.1)
.collect();
block.ffn.fc2.set_weight(&fc2_weight);
}
let n_vocab = config.n_vocab as usize;
let emb_data: Vec<f32> = (0..n_vocab * d_model)
.map(|i| (i as f32 * 0.001).sin() * 0.1)
.collect();
decoder.token_embedding_mut().copy_from_slice(&emb_data);
decoder.finalize_weights();
decoder
}
#[test]
fn test_decoder_block_new() {
let block = DecoderBlock::new(64, 4, 256);
assert_eq!(block.self_attn.d_model(), 64);
assert_eq!(block.cross_attn.d_model(), 64);
assert_eq!(block.ffn.d_model, 64);
}
#[test]
fn test_decoder_block_forward() {
let block = DecoderBlock::new(8, 2, 32);
let x = vec![0.1_f32; 16]; let encoder_out = vec![0.1_f32; 24];
let output = block
.forward(&x, &encoder_out, None)
.expect("forward should succeed");
assert_eq!(output.len(), 16); }
#[test]
fn test_decoder_block_with_causal_mask() {
let block = DecoderBlock::new(8, 2, 32);
let x = vec![0.1_f32; 16]; let encoder_out = vec![0.1_f32; 8]; let causal_mask = MultiHeadAttention::causal_mask(2);
let output = block
.forward(&x, &encoder_out, Some(&causal_mask))
.expect("forward should succeed");
assert_eq!(output.len(), 16);
}
#[test]
fn test_decoder_block_residual() {
let block = DecoderBlock::new(8, 2, 32);
let x = vec![1.0_f32; 8]; let encoder_out = vec![0.0_f32; 8];
let output = block
.forward(&x, &encoder_out, None)
.expect("forward should succeed");
assert_eq!(output.len(), 8);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_new() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
assert_eq!(decoder.n_layers(), 4);
assert_eq!(decoder.d_model(), 384);
assert_eq!(decoder.n_heads(), 6);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_vocab_size() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
assert_eq!(decoder.n_vocab(), 51865);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_max_len() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
assert_eq!(decoder.max_len(), 448);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_embedding_shapes() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
assert_eq!(
decoder.token_embedding().len(),
decoder.n_vocab() * decoder.d_model()
);
assert_eq!(
decoder.positional_embedding().len(),
decoder.max_len() * decoder.d_model()
);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_blocks_count() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
assert_eq!(decoder.blocks().len(), 4);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_embed_tokens_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![0, 1, 2];
let embeddings = decoder.embed_tokens(&tokens).expect("should succeed");
assert_eq!(embeddings.len(), 3 * decoder.d_model());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_embed_tokens_single() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![100];
let embeddings = decoder.embed_tokens(&tokens).expect("should succeed");
assert_eq!(embeddings.len(), decoder.d_model());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_embed_tokens_invalid() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![100000]; let result = decoder.embed_tokens(&tokens);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![0, 1, 2]; let encoder_out = vec![0.0_f32; 10 * 384];
let logits = decoder
.forward(&tokens, &encoder_out)
.expect("forward should succeed");
assert_eq!(logits.len(), 3 * decoder.n_vocab()); }
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_single_token() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![50258]; let encoder_out = vec![0.0_f32; 5 * 384];
let logits = decoder
.forward(&tokens, &encoder_out)
.expect("forward should succeed");
assert_eq!(logits.len(), decoder.n_vocab());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_empty_tokens() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens: Vec<u32> = vec![];
let encoder_out = vec![0.0_f32; 5 * 384];
let result = decoder.forward(&tokens, &encoder_out);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_sequence_too_long() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens: Vec<u32> = vec![0; 500]; let encoder_out = vec![0.0_f32; 5 * 384];
let result = decoder.forward(&tokens, &encoder_out);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_encoder_size_mismatch() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![0, 1];
let encoder_out = vec![0.0_f32; 100];
let result = decoder.forward(&tokens, &encoder_out);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_project_to_vocab_shape() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let hidden = vec![0.0_f32; 2 * 384]; let logits = decoder.project_to_vocab(&hidden, 2);
assert_eq!(logits.len(), 2 * decoder.n_vocab());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_project_to_vocab_single() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let hidden = vec![0.0_f32; 384]; let logits = decoder.project_to_vocab(&hidden, 1);
assert_eq!(logits.len(), decoder.n_vocab());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_project_to_vocab_correctness() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
let d_model = config.n_text_state as usize;
for i in 0..d_model {
decoder.token_embedding_mut()[i] = 1.0;
}
decoder.finalize_weights();
let hidden: Vec<f32> = vec![0.5; d_model];
let logits = decoder.project_to_vocab(&hidden, 1);
let expected = 0.5 * d_model as f32;
assert!(
(logits[0] - expected).abs() < 1e-3,
"expected {}, got {}",
expected,
logits[0]
);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_token_embedding_mut() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
decoder.token_embedding_mut()[0] = 1.0;
assert!((decoder.token_embedding()[0] - 1.0).abs() < 1e-6);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_positional_embedding_mut() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
decoder.positional_embedding_mut()[0] = 2.0;
assert!((decoder.positional_embedding()[0] - 2.0).abs() < 1e-6);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_output_different_for_different_tokens() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
let d_model = decoder.d_model();
for d in 0..d_model {
decoder.token_embedding_mut()[d] = 0.1;
decoder.token_embedding_mut()[d_model + d] = 0.2;
}
decoder.finalize_weights();
let encoder_out = vec![0.0_f32; 384];
let logits0 = decoder.forward(&[0], &encoder_out).expect("should succeed");
let logits1 = decoder.forward(&[1], &encoder_out).expect("should succeed");
let diff: f32 = logits0
.iter()
.zip(logits1.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(
diff > 0.0,
"Different tokens should produce different outputs"
);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_base_config() {
let config = ModelConfig::base();
let decoder = Decoder::new(&config);
assert_eq!(decoder.n_layers(), 6);
assert_eq!(decoder.d_model(), 512);
assert_eq!(decoder.n_heads(), 8);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_blocks_mut() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
let blocks = decoder.blocks_mut();
assert_eq!(blocks.len(), 4);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_ln_post() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let ln = decoder.ln_post();
assert_eq!(ln.normalized_shape, decoder.d_model());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_ln_post_mut() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
decoder.ln_post_mut().weight[0] = 3.0;
assert!((decoder.ln_post().weight[0] - 3.0).abs() < f32::EPSILON);
}
#[test]
fn test_layer_kv_cache_new() {
let cache = LayerKVCache::new(64, 100);
assert_eq!(cache.d_model, 64);
assert_eq!(cache.max_len, 100);
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
}
#[test]
fn test_layer_kv_cache_append() {
let mut cache = LayerKVCache::new(8, 100);
let key = vec![1.0_f32; 8]; let value = vec![2.0_f32; 8];
cache.append(&key, &value).expect("append should succeed");
assert_eq!(cache.len(), 1);
assert!(!cache.is_empty());
assert_eq!(cache.get_key().len(), 8);
assert_eq!(cache.get_value().len(), 8);
}
#[test]
fn test_layer_kv_cache_append_multiple() {
let mut cache = LayerKVCache::new(8, 100);
for _ in 0..3 {
let key = vec![1.0_f32; 8];
let value = vec![2.0_f32; 8];
cache.append(&key, &value).expect("append should succeed");
}
assert_eq!(cache.len(), 3);
assert_eq!(cache.get_key().len(), 24);
assert_eq!(cache.get_value().len(), 24);
}
#[test]
fn test_layer_kv_cache_overflow() {
let mut cache = LayerKVCache::new(8, 2);
cache.append(&[1.0; 8], &[2.0; 8]).expect("first append");
cache.append(&[1.0; 8], &[2.0; 8]).expect("second append");
let result = cache.append(&[1.0; 8], &[2.0; 8]);
assert!(result.is_err());
}
#[test]
fn test_layer_kv_cache_size_mismatch() {
let mut cache = LayerKVCache::new(8, 100);
let result = cache.append(&[1.0; 8], &[2.0; 16]);
assert!(result.is_err());
}
#[test]
fn test_layer_kv_cache_clear() {
let mut cache = LayerKVCache::new(8, 100);
cache.append(&[1.0; 8], &[2.0; 8]).expect("append");
cache.clear();
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
}
#[test]
fn test_layer_kv_cache_preallocated() {
let cache = LayerKVCache::new_preallocated(64, 100);
assert_eq!(cache.d_model, 64);
assert_eq!(cache.max_len, 100);
assert!(cache.is_empty());
assert!(cache.key.capacity() >= 64 * 100);
assert!(cache.value.capacity() >= 64 * 100);
}
#[test]
fn test_layer_kv_cache_remaining_capacity() {
let mut cache = LayerKVCache::new(8, 10);
assert_eq!(cache.remaining_capacity(), 10);
cache.append(&[1.0; 8], &[2.0; 8]).expect("append");
assert_eq!(cache.remaining_capacity(), 9);
cache
.append(&[1.0; 16], &[2.0; 16])
.expect("append 2 positions");
assert_eq!(cache.remaining_capacity(), 7);
}
#[test]
fn test_layer_kv_cache_is_full() {
let mut cache = LayerKVCache::new(8, 2);
assert!(!cache.is_full());
cache.append(&[1.0; 8], &[2.0; 8]).expect("append 1");
assert!(!cache.is_full());
cache.append(&[1.0; 8], &[2.0; 8]).expect("append 2");
assert!(cache.is_full());
}
#[test]
fn test_layer_kv_cache_reset_preserves_capacity() {
let mut cache = LayerKVCache::new(8, 100);
for _ in 0..10 {
cache.append(&[1.0; 8], &[2.0; 8]).expect("append");
}
let cap_before = cache.key.capacity();
assert!(cap_before >= 80);
cache.reset();
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
assert!(cache.key.capacity() >= cap_before);
}
#[test]
fn test_layer_kv_cache_append_batch() {
let mut cache = LayerKVCache::new(8, 100);
let keys = vec![1.0_f32; 24]; let values = vec![2.0_f32; 24];
cache.append_batch(&keys, &values, 3).expect("batch append");
assert_eq!(cache.len(), 3);
assert_eq!(cache.get_key().len(), 24);
assert_eq!(cache.get_value().len(), 24);
}
#[test]
fn test_layer_kv_cache_append_batch_mismatch() {
let mut cache = LayerKVCache::new(8, 100);
let keys = vec![1.0_f32; 16]; let values = vec![2.0_f32; 24];
let result = cache.append_batch(&keys, &values, 3);
assert!(result.is_err());
}
#[test]
fn test_layer_kv_cache_append_batch_overflow() {
let mut cache = LayerKVCache::new(8, 2);
let keys = vec![1.0_f32; 24];
let values = vec![2.0_f32; 24];
let result = cache.append_batch(&keys, &values, 3);
assert!(result.is_err());
}
#[test]
fn test_layer_kv_cache_get_key_range() {
let mut cache = LayerKVCache::new(4, 100);
for i in 0..5 {
let keys: Vec<f32> = (0..4).map(|d| (i * 4 + d) as f32).collect();
let values: Vec<f32> = (0..4).map(|d| (i * 4 + d + 100) as f32).collect();
cache.append(&keys, &values).expect("append");
}
let range = cache.get_key_range(1, 3).expect("should get range");
assert_eq!(range.len(), 8); assert!((range[0] - 4.0).abs() < f32::EPSILON); }
#[test]
fn test_layer_kv_cache_get_value_range() {
let mut cache = LayerKVCache::new(4, 100);
for i in 0..5 {
let keys: Vec<f32> = (0..4).map(|d| (i * 4 + d) as f32).collect();
let values: Vec<f32> = (0..4).map(|d| (i * 4 + d + 100) as f32).collect();
cache.append(&keys, &values).expect("append");
}
let range = cache.get_value_range(2, 4).expect("should get range");
assert_eq!(range.len(), 8);
assert!((range[0] - 108.0).abs() < f32::EPSILON); }
#[test]
fn test_layer_kv_cache_get_range_out_of_bounds() {
let mut cache = LayerKVCache::new(4, 100);
cache.append(&[1.0; 4], &[2.0; 4]).expect("append");
assert!(cache.get_key_range(0, 5).is_none());
assert!(cache.get_value_range(2, 3).is_none());
assert!(cache.get_key_range(3, 1).is_none());
}
#[test]
fn test_layer_kv_cache_memory_bytes() {
let mut cache = LayerKVCache::new(8, 100);
assert_eq!(cache.memory_bytes(), 0);
cache.append(&[1.0; 8], &[2.0; 8]).expect("append");
assert_eq!(cache.memory_bytes(), 64);
}
#[test]
fn test_layer_kv_cache_capacity_bytes() {
let cache = LayerKVCache::new_preallocated(8, 10);
assert!(cache.capacity_bytes() >= 640);
}
#[test]
fn test_transposed_cache_new() {
let cache = LayerKVCacheTransposed::new(64, 100);
assert_eq!(cache.d_model, 64);
assert_eq!(cache.max_len, 100);
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
}
#[test]
fn test_transposed_cache_append() {
let mut cache = LayerKVCacheTransposed::new(4, 100);
let key = vec![1.0_f32, 2.0, 3.0, 4.0];
let value = vec![10.0_f32, 20.0, 30.0, 40.0];
cache.append(&key, &value).expect("append should succeed");
assert_eq!(cache.len(), 1);
assert!(!cache.is_empty());
assert_eq!(cache.get_key().len(), 4);
assert_eq!(cache.get_value_transposed().len(), 4);
let v_t = cache.get_value_transposed();
assert_eq!(v_t, &[10.0_f32, 20.0, 30.0, 40.0]);
}
#[test]
fn test_transposed_cache_append_multiple() {
let mut cache = LayerKVCacheTransposed::new(4, 100);
let key1 = vec![1.0_f32, 2.0, 3.0, 4.0];
let value1 = vec![10.0_f32, 20.0, 30.0, 40.0];
cache.append(&key1, &value1).expect("first append");
let key2 = vec![5.0_f32, 6.0, 7.0, 8.0];
let value2 = vec![50.0_f32, 60.0, 70.0, 80.0];
cache.append(&key2, &value2).expect("second append");
assert_eq!(cache.len(), 2);
assert_eq!(cache.get_key().len(), 8);
assert_eq!(cache.get_value_transposed().len(), 8);
let v_t = cache.get_value_transposed();
assert_eq!(v_t, &[10.0_f32, 50.0, 20.0, 60.0, 30.0, 70.0, 40.0, 80.0]);
}
#[test]
fn test_transposed_cache_get_feature() {
let mut cache = LayerKVCacheTransposed::new(4, 100);
cache
.append(&[1.0; 4], &[10.0_f32, 20.0, 30.0, 40.0])
.expect("append 1");
cache
.append(&[2.0; 4], &[11.0_f32, 21.0, 31.0, 41.0])
.expect("append 2");
cache
.append(&[3.0; 4], &[12.0_f32, 22.0, 32.0, 42.0])
.expect("append 3");
let f0 = cache.get_value_feature(0).expect("feature 0");
assert_eq!(f0, &[10.0_f32, 11.0, 12.0]);
let f1 = cache.get_value_feature(1).expect("feature 1");
assert_eq!(f1, &[20.0_f32, 21.0, 22.0]);
let f3 = cache.get_value_feature(3).expect("feature 3");
assert_eq!(f3, &[40.0_f32, 41.0, 42.0]);
assert!(cache.get_value_feature(4).is_none());
}
#[test]
fn test_transposed_cache_apply_attention() {
let mut cache = LayerKVCacheTransposed::new(2, 100);
cache.append(&[0.0; 2], &[1.0_f32, 2.0]).expect("pos 0");
cache.append(&[0.0; 2], &[3.0_f32, 4.0]).expect("pos 1");
let scores = vec![0.5_f32, 0.5];
let output = cache.apply_attention(&scores, 1);
assert_eq!(output.len(), 2);
assert!((output[0] - 2.0).abs() < 1e-6);
assert!((output[1] - 3.0).abs() < 1e-6);
}
#[test]
fn test_transposed_cache_clear() {
let mut cache = LayerKVCacheTransposed::new(4, 100);
cache.append(&[1.0; 4], &[2.0; 4]).expect("append");
cache.clear();
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
assert_eq!(cache.memory_bytes(), 0);
}
#[test]
fn test_transposed_cache_memory_bytes() {
let mut cache = LayerKVCacheTransposed::new(4, 100);
assert_eq!(cache.memory_bytes(), 0);
cache.append(&[1.0; 4], &[2.0; 4]).expect("append");
assert_eq!(cache.memory_bytes(), 32);
}
#[test]
fn test_circular_kv_buffer_new() {
let buffer = CircularKVBuffer::new(10, 8);
assert_eq!(buffer.len(), 0);
assert!(buffer.is_empty());
assert!(!buffer.is_full());
assert_eq!(buffer.window_size(), 10);
}
#[test]
fn test_circular_kv_buffer_append() {
let mut buffer = CircularKVBuffer::new(5, 4);
let key = vec![1.0_f32; 4];
let value = vec![2.0_f32; 4];
buffer.append(&key, &value);
assert_eq!(buffer.len(), 1);
assert!(!buffer.is_empty());
assert!(!buffer.is_full());
}
#[test]
fn test_circular_kv_buffer_fill() {
let mut buffer = CircularKVBuffer::new(3, 2);
for i in 0..3 {
let key = vec![i as f32; 2];
let value = vec![(i + 10) as f32; 2];
buffer.append(&key, &value);
}
assert_eq!(buffer.len(), 3);
assert!(buffer.is_full());
let keys = buffer.get_keys_linear();
let values = buffer.get_values_linear();
assert_eq!(keys.len(), 6);
assert_eq!(values.len(), 6);
assert_eq!(keys[0..2], [0.0, 0.0]);
assert_eq!(keys[2..4], [1.0, 1.0]);
assert_eq!(keys[4..6], [2.0, 2.0]);
}
#[test]
fn test_circular_kv_buffer_wrap_around() {
let mut buffer = CircularKVBuffer::new(3, 2);
for i in 0..3 {
let key = vec![i as f32; 2];
let value = vec![(i + 10) as f32; 2];
buffer.append(&key, &value);
}
buffer.append(&[99.0, 99.0], &[199.0, 199.0]);
assert_eq!(buffer.len(), 3); assert!(buffer.is_full());
let keys = buffer.get_keys_linear();
assert_eq!(keys[0..2], [1.0, 1.0]);
assert_eq!(keys[2..4], [2.0, 2.0]);
assert_eq!(keys[4..6], [99.0, 99.0]);
}
#[test]
fn test_circular_kv_buffer_batch_append() {
let mut buffer = CircularKVBuffer::new(10, 4);
let keys = vec![1.0_f32; 12]; let values = vec![2.0_f32; 12];
buffer.append_batch(&keys, &values, 3);
assert_eq!(buffer.len(), 3);
}
#[test]
fn test_circular_kv_buffer_reset() {
let mut buffer = CircularKVBuffer::new(5, 4);
buffer.append(&[1.0; 4], &[2.0; 4]);
buffer.append(&[3.0; 4], &[4.0; 4]);
assert_eq!(buffer.len(), 2);
buffer.reset();
assert_eq!(buffer.len(), 0);
assert!(buffer.is_empty());
}
#[test]
fn test_circular_kv_buffer_memory_bytes() {
let buffer = CircularKVBuffer::new(10, 8);
assert_eq!(buffer.memory_bytes(), 640);
}
#[test]
fn test_decoder_kv_cache_new() {
let cache = DecoderKVCache::new(4, 64, 100);
assert_eq!(cache.n_layers, 4);
assert_eq!(cache.d_model, 64);
assert_eq!(cache.max_len, 100);
assert!(cache.is_empty());
assert_eq!(cache.seq_len(), 0);
}
#[test]
fn test_decoder_kv_cache_layer_count() {
let cache = DecoderKVCache::new(4, 64, 100);
assert_eq!(cache.self_attn_cache.len(), 4);
assert_eq!(cache.cross_attn_cache.len(), 4);
}
#[test]
fn test_decoder_kv_cache_clear() {
let mut cache = DecoderKVCache::new(4, 8, 100);
cache.self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.expect("append");
cache.cross_attn_cached = true;
cache.clear();
assert!(cache.is_empty());
assert!(!cache.cross_attn_cached);
}
#[test]
fn test_decoder_kv_cache_clear_self_attn() {
let mut cache = DecoderKVCache::new(4, 8, 100);
cache.self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.expect("append");
cache.cross_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.expect("append");
cache.cross_attn_cached = true;
cache.clear_self_attn();
assert!(cache.self_attn_cache[0].is_empty());
assert!(!cache.cross_attn_cache[0].is_empty()); assert!(cache.cross_attn_cached); }
#[test]
fn test_decoder_kv_cache_memory_bytes() {
let mut cache = DecoderKVCache::new(2, 8, 100);
assert_eq!(cache.memory_bytes(), 0);
cache.self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.expect("append");
assert_eq!(cache.memory_bytes(), 64);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_create_kv_cache() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let cache = decoder.create_kv_cache();
assert_eq!(cache.n_layers, decoder.n_layers());
assert_eq!(cache.d_model, decoder.d_model());
assert_eq!(cache.max_len, decoder.max_len());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_one_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_kv_cache();
let encoder_out = vec![0.0_f32; 5 * 384];
let logits = decoder
.forward_one(0, &encoder_out, &mut cache)
.expect("forward_one should succeed");
assert_eq!(logits.len(), decoder.n_vocab());
assert_eq!(cache.seq_len(), 1);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_one_multiple() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_kv_cache();
let encoder_out = vec![0.0_f32; 5 * 384];
for token in 0..3 {
let _ = decoder
.forward_one(token, &encoder_out, &mut cache)
.expect("forward_one should succeed");
}
assert_eq!(cache.seq_len(), 3);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_one_invalid_token() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_kv_cache();
let encoder_out = vec![0.0_f32; 5 * 384];
let result = decoder.forward_one(100000, &encoder_out, &mut cache);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_generate_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let encoder_out = vec![0.0_f32; 5 * 384];
let initial = vec![50258_u32]; let eos = 50257_u32;
let tokens = decoder
.generate(&encoder_out, &initial, 5, eos)
.expect("generate should succeed");
assert!(tokens.len() >= initial.len());
assert!(tokens.len() <= 5);
assert_eq!(tokens[0], 50258);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_generate_stops_at_eos() {
let config = ModelConfig::tiny();
let mut decoder = Decoder::new(&config);
let d_model = decoder.d_model();
for d in 0..d_model {
decoder.token_embedding_mut()[d_model + d] = if d == 0 { 10.0 } else { 0.0 };
}
let eos_start = 50257 * d_model;
for d in 0..d_model {
decoder.token_embedding_mut()[eos_start + d] = if d == 0 { 10.0 } else { 0.0 };
}
decoder.finalize_weights();
let encoder_out = vec![0.0_f32; 384];
let initial = vec![1_u32];
let eos = 50257_u32;
let tokens = decoder
.generate(&encoder_out, &initial, 100, eos)
.expect("generate should succeed");
assert!(tokens.len() < 100);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_kv_cache_reuse() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_kv_cache();
let encoder_out = vec![0.0_f32; 5 * 384];
let logits1 = decoder
.forward_one(0, &encoder_out, &mut cache)
.expect("first forward");
let logits2 = decoder
.forward_one(1, &encoder_out, &mut cache)
.expect("second forward");
assert_eq!(logits1.len(), decoder.n_vocab());
assert_eq!(logits2.len(), decoder.n_vocab());
assert_eq!(cache.seq_len(), 2);
}
#[test]
fn test_batch_decoder_cache_new() {
let cache = BatchDecoderCache::new(3, 4, 64, 100);
assert_eq!(cache.batch_size(), 3);
assert_eq!(cache.n_layers, 4);
assert_eq!(cache.d_model, 64);
assert!(cache.is_empty());
}
#[test]
fn test_batch_decoder_cache_get_cache() {
let cache = BatchDecoderCache::new(3, 4, 64, 100);
let item0 = cache.get_cache(0);
assert!(item0.is_some());
assert_eq!(item0.unwrap().n_layers, 4);
let item3 = cache.get_cache(3);
assert!(item3.is_none()); }
#[test]
fn test_batch_decoder_cache_get_cache_mut() {
let mut cache = BatchDecoderCache::new(2, 4, 8, 100);
{
let item0 = cache.get_cache_mut(0).unwrap();
item0.self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
}
assert_eq!(cache.get_cache(0).unwrap().seq_len(), 1);
assert_eq!(cache.get_cache(1).unwrap().seq_len(), 0);
}
#[test]
fn test_batch_decoder_cache_clear_all() {
let mut cache = BatchDecoderCache::new(2, 4, 8, 100);
cache.get_cache_mut(0).unwrap().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
cache.get_cache_mut(1).unwrap().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
cache.clear_all();
assert!(cache.get_cache(0).unwrap().is_empty());
assert!(cache.get_cache(1).unwrap().is_empty());
}
#[test]
fn test_batch_decoder_cache_seq_lengths() {
let mut cache = BatchDecoderCache::new(3, 4, 8, 100);
cache.get_cache_mut(0).unwrap().self_attn_cache[0]
.append(&[1.0; 16], &[2.0; 16])
.unwrap(); cache.get_cache_mut(1).unwrap().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
let lengths = cache.seq_lengths();
assert_eq!(lengths, vec![2, 1, 0]);
}
#[test]
fn test_batch_decoder_cache_max_seq_len() {
let mut cache = BatchDecoderCache::new(3, 4, 8, 100);
cache.get_cache_mut(0).unwrap().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
cache.get_cache_mut(1).unwrap().self_attn_cache[0]
.append(&[1.0; 24], &[2.0; 24])
.unwrap();
assert_eq!(cache.max_seq_len(), 3);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_create_batch_cache() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let cache = decoder.create_batch_cache(4);
assert_eq!(cache.batch_size(), 4);
assert_eq!(cache.n_layers, decoder.n_layers());
assert_eq!(cache.d_model, decoder.d_model());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_batch_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens_batch = vec![
vec![0_u32, 1, 2], vec![3_u32, 4], ];
let encoder_outputs = vec![
vec![0.0_f32; 5 * 384], vec![0.0_f32; 3 * 384], ];
let result = decoder
.forward_batch(&tokens_batch, &encoder_outputs)
.expect("forward_batch should succeed");
assert_eq!(result.batch_size(), 2);
assert_eq!(result.logits.len(), 2);
assert_eq!(result.logits[0].len(), 3 * decoder.n_vocab()); assert_eq!(result.logits[1].len(), 2 * decoder.n_vocab());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_batch_empty() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let result = decoder.forward_batch(&[], &[]);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_batch_mismatch() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let tokens = vec![vec![0_u32]];
let encoders = vec![vec![0.0_f32; 384], vec![0.0_f32; 384]];
let result = decoder.forward_batch(&tokens, &encoders);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_one_batch_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_batch_cache(2);
let tokens = vec![0_u32, 1_u32]; let encoder_outputs = vec![vec![0.0_f32; 5 * 384], vec![0.0_f32; 3 * 384]];
let result = decoder
.forward_one_batch(&tokens, &encoder_outputs, &mut cache)
.expect("forward_one_batch should succeed");
assert_eq!(result.batch_size(), 2);
assert_eq!(result.logits.len(), 2);
assert_eq!(result.logits[0].len(), decoder.n_vocab());
assert_eq!(result.logits[1].len(), decoder.n_vocab());
assert_eq!(cache.get_cache(0).unwrap().seq_len(), 1);
assert_eq!(cache.get_cache(1).unwrap().seq_len(), 1);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_one_batch_multiple_steps() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_batch_cache(2);
let encoder_outputs = vec![vec![0.0_f32; 5 * 384], vec![0.0_f32; 5 * 384]];
decoder
.forward_one_batch(&[0, 1], &encoder_outputs, &mut cache)
.unwrap();
decoder
.forward_one_batch(&[2, 3], &encoder_outputs, &mut cache)
.unwrap();
decoder
.forward_one_batch(&[4, 5], &encoder_outputs, &mut cache)
.unwrap();
assert_eq!(cache.get_cache(0).unwrap().seq_len(), 3);
assert_eq!(cache.get_cache(1).unwrap().seq_len(), 3);
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_forward_one_batch_size_mismatch() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let mut cache = decoder.create_batch_cache(2);
let tokens = vec![0_u32, 1, 2]; let encoder_outputs = vec![vec![0.0_f32; 384], vec![0.0_f32; 384]];
let result = decoder.forward_one_batch(&tokens, &encoder_outputs, &mut cache);
assert!(result.is_err());
}
#[test]
#[ignore = "Allocates large model - run with --ignored"]
fn test_decoder_generate_batch_basic() {
let config = ModelConfig::tiny();
let decoder = Decoder::new(&config);
let encoder_outputs = vec![vec![0.0_f32; 5 * 384], vec![0.0_f32; 5 * 384]];
let initial_tokens = vec![
vec![50258_u32], vec![50258_u32], ];
let eos = 50257_u32;
let result = decoder
.generate_batch(&encoder_outputs, &initial_tokens, 5, eos)
.expect("generate_batch should succeed");
assert_eq!(result.len(), 2);
assert!(result[0].len() >= 1);
assert!(result[1].len() >= 1);
}
#[test]
fn test_batch_decoder_output_batch_size() {
let output = BatchDecoderOutput {
logits: vec![vec![0.0; 100], vec![0.0; 100], vec![0.0; 100]],
seq_lengths: vec![5, 3, 4],
};
assert_eq!(output.batch_size(), 3);
}
#[test]
fn test_batch_decoder_output_get_logits() {
let output = BatchDecoderOutput {
logits: vec![vec![1.0; 10], vec![2.0; 10]],
seq_lengths: vec![1, 1],
};
assert_eq!(output.get_logits(0).unwrap()[0], 1.0);
assert_eq!(output.get_logits(1).unwrap()[0], 2.0);
assert!(output.get_logits(2).is_none());
}
#[test]
fn test_batch_decoder_output_is_empty() {
let empty = BatchDecoderOutput {
logits: vec![],
seq_lengths: vec![],
};
assert!(empty.is_empty());
let non_empty = BatchDecoderOutput {
logits: vec![vec![0.0]],
seq_lengths: vec![1],
};
assert!(!non_empty.is_empty());
}
#[test]
fn test_streaming_kv_cache_new() {
let cache = StreamingKVCache::new(4, 64, 100, 20);
assert_eq!(cache.window_size(), 100);
assert_eq!(cache.context_overlap(), 20);
assert!(cache.is_empty());
assert_eq!(cache.seq_len(), 0);
assert_eq!(cache.total_tokens(), 0);
assert_eq!(cache.slide_count(), 0);
}
#[test]
fn test_streaming_kv_cache_low_latency() {
let cache = StreamingKVCache::low_latency(4, 64);
assert_eq!(cache.window_size(), 64);
assert_eq!(cache.context_overlap(), 16);
}
#[test]
fn test_streaming_kv_cache_ultra_low_latency() {
let cache = StreamingKVCache::ultra_low_latency(4, 64);
assert_eq!(cache.window_size(), 32);
assert_eq!(cache.context_overlap(), 8);
}
#[test]
fn test_streaming_kv_cache_standard() {
let cache = StreamingKVCache::standard(4, 64);
assert_eq!(cache.window_size(), 448);
assert_eq!(cache.context_overlap(), 64);
}
#[test]
fn test_streaming_kv_cache_overlap_clamped() {
let cache = StreamingKVCache::new(4, 64, 100, 80);
assert_eq!(cache.context_overlap(), 50); }
#[test]
fn test_streaming_kv_cache_remaining_capacity() {
let mut cache = StreamingKVCache::new(4, 8, 10, 2);
assert_eq!(cache.remaining_capacity(), 10);
cache.inner_mut().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
assert_eq!(cache.remaining_capacity(), 9);
}
#[test]
fn test_streaming_kv_cache_will_slide() {
let mut cache = StreamingKVCache::new(4, 8, 3, 1);
assert!(!cache.will_slide());
cache.inner_mut().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
cache.inner_mut().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
cache.inner_mut().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
assert!(cache.will_slide());
}
#[test]
fn test_streaming_kv_cache_append_with_slide() {
let mut cache = StreamingKVCache::new(2, 8, 4, 2);
for _ in 0..3 {
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
}
assert_eq!(cache.seq_len(), 3);
assert_eq!(cache.slide_count(), 0);
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
assert!(cache.slide_count() > 0);
assert!(cache.seq_len() <= cache.window_size());
}
#[test]
fn test_streaming_kv_cache_slide_preserves_overlap() {
let mut cache = StreamingKVCache::new(1, 4, 5, 2);
for i in 0..5 {
let keys: Vec<f32> = (0..4).map(|d| (i * 4 + d) as f32).collect();
let values: Vec<f32> = (0..4).map(|d| (i * 4 + d + 100) as f32).collect();
cache.inner_mut().self_attn_cache[0]
.append(&keys, &values)
.unwrap();
}
assert_eq!(cache.seq_len(), 5);
cache.slide_window().unwrap();
assert_eq!(cache.seq_len(), 2);
assert_eq!(cache.slide_count(), 1);
let keys = cache.inner().self_attn_cache[0].get_key();
assert!((keys[0] - 12.0).abs() < 0.01);
}
#[test]
fn test_streaming_kv_cache_reset() {
let mut cache = StreamingKVCache::new(2, 8, 10, 2);
for _ in 0..15 {
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
}
let prev_total = cache.total_tokens();
let prev_slides = cache.slide_count();
cache.reset();
assert!(cache.is_empty());
assert_eq!(cache.total_tokens(), prev_total);
assert_eq!(cache.slide_count(), prev_slides);
}
#[test]
fn test_streaming_kv_cache_full_reset() {
let mut cache = StreamingKVCache::new(2, 8, 10, 2);
for _ in 0..15 {
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
}
cache.full_reset();
assert!(cache.is_empty());
assert_eq!(cache.total_tokens(), 0);
assert_eq!(cache.slide_count(), 0);
}
#[test]
fn test_streaming_kv_cache_warm_up() {
let mut cache = StreamingKVCache::new(2, 4, 10, 3);
let keys: Vec<f32> = (0..20).map(|i| i as f32).collect();
let values: Vec<f32> = (0..20).map(|i| i as f32 + 100.0).collect();
cache.warm_up(0, &keys, &values).unwrap();
assert_eq!(cache.seq_len(), 3);
let cached_keys = cache.inner().self_attn_cache[0].get_key();
assert!((cached_keys[0] - 8.0).abs() < 0.01);
}
#[test]
fn test_streaming_kv_cache_warm_up_invalid_layer() {
let mut cache = StreamingKVCache::new(2, 4, 10, 3);
let result = cache.warm_up(5, &[1.0; 8], &[2.0; 8]); assert!(result.is_err());
}
#[test]
fn test_streaming_kv_cache_stats() {
let mut cache = StreamingKVCache::new(2, 8, 10, 2);
for _ in 0..5 {
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
}
let stats = cache.stats();
assert_eq!(stats.seq_len, 5);
assert_eq!(stats.total_tokens, 5);
assert_eq!(stats.window_size, 10);
assert_eq!(stats.context_overlap, 2);
assert!((stats.utilization() - 0.5).abs() < 0.01); }
#[test]
fn test_streaming_cache_stats_utilization() {
let stats = StreamingCacheStats {
seq_len: 25,
total_tokens: 100,
slide_count: 3,
window_size: 50,
context_overlap: 10,
memory_bytes: 1000,
};
assert!((stats.utilization() - 0.5).abs() < 0.01); }
#[test]
fn test_streaming_cache_stats_tokens_per_slide() {
let stats = StreamingCacheStats {
seq_len: 25,
total_tokens: 100,
slide_count: 4,
window_size: 50,
context_overlap: 10,
memory_bytes: 1000,
};
assert!((stats.tokens_per_slide() - 25.0).abs() < 0.01); }
#[test]
fn test_streaming_cache_stats_tokens_per_slide_no_slides() {
let stats = StreamingCacheStats {
seq_len: 10,
total_tokens: 10,
slide_count: 0,
window_size: 50,
context_overlap: 10,
memory_bytes: 1000,
};
assert!((stats.tokens_per_slide() - 10.0).abs() < 0.01); }
#[test]
fn test_streaming_cache_stats_zero_window() {
let stats = StreamingCacheStats {
seq_len: 0,
total_tokens: 0,
slide_count: 0,
window_size: 0,
context_overlap: 0,
memory_bytes: 0,
};
assert!((stats.utilization() - 0.0).abs() < f32::EPSILON);
}
#[test]
fn test_streaming_kv_cache_inner_accessors() {
let mut cache = StreamingKVCache::new(2, 8, 10, 2);
let inner_ref = cache.inner();
assert_eq!(inner_ref.n_layers, 2);
let inner_mut = cache.inner_mut();
inner_mut.self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
assert_eq!(cache.seq_len(), 1);
}
#[test]
fn test_streaming_kv_cache_memory_bytes() {
let mut cache = StreamingKVCache::new(2, 8, 10, 2);
assert_eq!(cache.memory_bytes(), 0);
cache.inner_mut().self_attn_cache[0]
.append(&[1.0; 8], &[2.0; 8])
.unwrap();
assert!(cache.memory_bytes() > 0);
}
#[test]
fn test_streaming_kv_cache_continuous_streaming() {
let mut cache = StreamingKVCache::new(2, 8, 20, 5);
for _ in 0..100 {
cache.append_with_slide(0, &[1.0; 8], &[2.0; 8]).unwrap();
}
assert!(cache.seq_len() <= cache.window_size());
assert!(cache.slide_count() > 0);
assert_eq!(cache.total_tokens(), 100);
let stats = cache.stats();
assert!(stats.tokens_per_slide() > 0.0);
}
#[test]
fn test_cross_attention_uses_encoder_output() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let tokens = vec![0_u32, 1, 2];
let encoder_a = vec![1.0_f32; 5 * config.n_audio_state as usize];
let encoder_b = vec![-1.0_f32; 5 * config.n_audio_state as usize];
let output_a = decoder.forward(&tokens, &encoder_a).unwrap();
let output_b = decoder.forward(&tokens, &encoder_b).unwrap();
let diff: f32 = output_a
.iter()
.zip(output_b.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(
diff > 0.001,
"Cross-attention not working: outputs are identical for different encoder inputs (diff={diff})"
);
}
#[test]
fn test_cross_attention_output_varies_with_encoder() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let tokens = vec![50257_u32];
let mut outputs = Vec::new();
for i in 0..3 {
let encoder_out: Vec<f32> = (0..5 * config.n_audio_state as usize)
.map(|j| ((i * 1000 + j) as f32).sin())
.collect();
let logits = decoder.forward(&tokens, &encoder_out).unwrap();
outputs.push(logits);
}
for i in 0..outputs.len() {
for j in (i + 1)..outputs.len() {
let diff: f32 = outputs[i]
.iter()
.zip(outputs[j].iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(
diff > 0.001,
"Cross-attention outputs {i} and {j} are too similar (diff={diff})"
);
}
}
}
#[cfg(feature = "realizar-inference")]
mod paged_kv_tests {
use super::*;
#[test]
fn test_paged_decoder_kv_cache_new() {
let config = ModelConfig::tiny();
let cache = PagedDecoderKVCache::new(&config, 64);
assert_eq!(cache.num_layers(), config.n_text_layer as usize);
assert_eq!(cache.total_pages(), 64);
assert_eq!(cache.used_pages(), 0);
}
#[test]
fn test_paged_decoder_kv_cache_allocate_sequence() {
let config = ModelConfig::tiny();
let mut cache = PagedDecoderKVCache::new(&config, 64);
let seq_id = cache.allocate_sequence(32).unwrap(); assert!(cache.used_pages() > 0);
assert!(cache.has_sequence(seq_id));
}
#[test]
fn test_paged_decoder_kv_cache_append_kv() {
let config = ModelConfig::tiny();
let mut cache = PagedDecoderKVCache::new(&config, 64);
let d_model = config.n_text_state as usize;
let n_layers = config.n_text_layer as usize;
let seq_id = cache.allocate_sequence(0).unwrap();
let key = vec![1.0_f32; d_model];
let value = vec![2.0_f32; d_model];
for layer in 0..n_layers {
cache.append(seq_id, layer, &key, &value).unwrap();
}
cache.increment_seq_len(seq_id);
assert_eq!(cache.seq_len(seq_id), 1);
}
#[test]
fn test_paged_decoder_kv_cache_read_kv() {
let config = ModelConfig::tiny();
let mut cache = PagedDecoderKVCache::new(&config, 64);
let d_model = config.n_text_state as usize;
let n_layers = config.n_text_layer as usize;
let seq_id = cache.allocate_sequence(0).unwrap();
let key: Vec<f32> = (0..d_model).map(|i| i as f32).collect();
let value: Vec<f32> = (0..d_model).map(|i| (i + 100) as f32).collect();
for layer in 0..n_layers {
cache.append(seq_id, layer, &key, &value).unwrap();
}
cache.increment_seq_len(seq_id);
let (read_key, read_value) = cache.get_kv(seq_id, 0).unwrap();
assert_eq!(read_key.len(), d_model);
assert_eq!(read_value.len(), d_model);
for i in 0..d_model {
assert!((read_key[i] - i as f32).abs() < 1e-5);
assert!((read_value[i] - (i + 100) as f32).abs() < 1e-5);
}
}
#[test]
fn test_paged_decoder_kv_cache_memory_efficiency() {
let config = ModelConfig::tiny();
let d_model = config.n_text_state as usize;
let n_layers = config.n_text_layer as usize;
let max_seq_len = 448;
let naive_bytes = 2 * n_layers * max_seq_len * d_model * 4;
let mut cache = PagedDecoderKVCache::new(&config, 64);
let seq_id = cache.allocate_sequence(0).unwrap();
for _ in 0..10 {
let key = vec![1.0_f32; d_model];
let value = vec![2.0_f32; d_model];
for layer in 0..n_layers {
cache.append(seq_id, layer, &key, &value).unwrap();
}
cache.increment_seq_len(seq_id);
}
let paged_bytes = cache.memory_bytes();
assert!(
paged_bytes < naive_bytes / 4,
"Paged cache should use <25% of naive allocation: {} vs {}",
paged_bytes,
naive_bytes
);
}
#[test]
fn test_paged_decoder_kv_cache_free_sequence() {
let config = ModelConfig::tiny();
let mut cache = PagedDecoderKVCache::new(&config, 64);
let seq_id = cache.allocate_sequence(32).unwrap();
let used_before = cache.used_pages();
assert!(used_before > 0);
cache.free_sequence(seq_id).unwrap();
assert_eq!(cache.used_pages(), 0);
assert!(!cache.has_sequence(seq_id));
}
#[test]
fn test_paged_decoder_kv_cache_multiple_sequences() {
let config = ModelConfig::tiny();
let mut cache = PagedDecoderKVCache::new(&config, 128);
let seq1 = cache.allocate_sequence(16).unwrap();
let seq2 = cache.allocate_sequence(16).unwrap();
let seq3 = cache.allocate_sequence(16).unwrap();
assert!(cache.has_sequence(seq1));
assert!(cache.has_sequence(seq2));
assert!(cache.has_sequence(seq3));
cache.free_sequence(seq2).unwrap();
assert!(!cache.has_sequence(seq2));
assert!(cache.has_sequence(seq1));
assert!(cache.has_sequence(seq3));
}
#[test]
fn test_paged_decoder_kv_cache_out_of_memory() {
let config = ModelConfig::tiny();
let mut cache = PagedDecoderKVCache::new(&config, 4);
let result = cache.allocate_sequence(1000);
assert!(result.is_err());
}
#[test]
fn test_paged_vs_baseline_numerical_equivalence() {
let config = ModelConfig::tiny();
let d_model = config.n_text_state as usize;
let n_layers = config.n_text_layer as usize;
let mut baselines: Vec<LayerKVCache> = (0..n_layers)
.map(|_| LayerKVCache::new(d_model, 100))
.collect();
let mut paged = PagedDecoderKVCache::new(&config, 64);
let seq_id = paged.allocate_sequence(0).unwrap();
for i in 0..10 {
let key: Vec<f32> = (0..d_model).map(|j| (i * d_model + j) as f32).collect();
let value: Vec<f32> = (0..d_model)
.map(|j| ((i * d_model + j) as f32).sin())
.collect();
for layer in 0..n_layers {
baselines[layer].append(&key, &value).unwrap();
paged.append(seq_id, layer, &key, &value).unwrap();
}
paged.increment_seq_len(seq_id);
}
for layer in 0..n_layers {
let baseline_key = baselines[layer].get_key();
let baseline_value = baselines[layer].get_value();
let (paged_keys, paged_values) = paged.get_all_kv(seq_id, layer).unwrap();
assert_eq!(
baseline_key.len(),
paged_keys.len(),
"Layer {layer} key length mismatch"
);
for i in 0..baseline_key.len() {
assert!(
(baseline_key[i] - paged_keys[i]).abs() < 1e-6,
"Layer {layer} key mismatch at {}: {} vs {}",
i,
baseline_key[i],
paged_keys[i]
);
}
for i in 0..baseline_value.len() {
assert!(
(baseline_value[i] - paged_values[i]).abs() < 1e-6,
"Layer {layer} value mismatch at {}: {} vs {}",
i,
baseline_value[i],
paged_values[i]
);
}
}
}
#[test]
fn test_decoder_create_paged_kv_cache() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let cache = decoder.create_paged_kv_cache(64);
assert_eq!(cache.num_layers(), config.n_text_layer as usize);
assert_eq!(cache.total_pages(), 64);
}
#[test]
fn test_decoder_forward_one_paged_matches_baseline() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let token = 50257_u32;
let mut baseline_cache = decoder.create_kv_cache();
let baseline_logits = decoder
.forward_one(token, &encoder_output, &mut baseline_cache)
.unwrap();
let mut paged_cache = decoder.create_paged_kv_cache(64);
let seq_id = paged_cache.allocate_sequence(0).unwrap();
let paged_logits = decoder
.forward_one_paged(token, &encoder_output, &mut paged_cache, seq_id)
.unwrap();
assert_eq!(baseline_logits.len(), paged_logits.len());
for i in 0..baseline_logits.len() {
assert!(
(baseline_logits[i] - paged_logits[i]).abs() < 1e-5,
"Logit mismatch at {}: {} vs {}",
i,
baseline_logits[i],
paged_logits[i]
);
}
}
#[test]
fn test_decoder_generate_paged_matches_baseline() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let initial_tokens = vec![50257_u32]; let eos_token = 50256_u32;
let baseline_tokens = decoder
.generate(&encoder_output, &initial_tokens, 10, eos_token)
.unwrap();
let paged_tokens = decoder
.generate_paged(&encoder_output, &initial_tokens, 10, eos_token)
.unwrap();
assert_eq!(
baseline_tokens, paged_tokens,
"Token mismatch: baseline {:?} vs paged {:?}",
baseline_tokens, paged_tokens
);
}
#[test]
fn test_decoder_generate_paged_memory_efficiency() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let n_layers = config.n_text_layer as usize;
let max_len = decoder.max_len;
let _baseline_cache = decoder.create_kv_cache();
let baseline_capacity = 2 * n_layers * max_len * d_model * 4;
let mut paged_cache = decoder.create_paged_kv_cache(64);
let seq_id = paged_cache.allocate_sequence(0).unwrap();
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
for _ in 0..5 {
decoder
.forward_one_paged(50257, &encoder_output, &mut paged_cache, seq_id)
.unwrap();
}
let paged_bytes = paged_cache.memory_bytes();
assert!(
paged_bytes < baseline_capacity / 4,
"Paged should use <25% memory: {} vs {} capacity",
paged_bytes,
baseline_capacity
);
}
#[test]
fn test_decoder_generate_paged_multi_sequence() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder1: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let encoder2: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).cos()).collect();
let mut paged_cache = decoder.create_paged_kv_cache(128);
let seq1 = paged_cache.allocate_sequence(0).unwrap();
let seq2 = paged_cache.allocate_sequence(0).unwrap();
let mut tokens1 = vec![50257_u32];
let mut tokens2 = vec![50257_u32];
for _ in 0..3 {
let logits1 = decoder
.forward_one_paged(*tokens1.last().unwrap(), &encoder1, &mut paged_cache, seq1)
.unwrap();
let next1 = logits1
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
.map_or(50256, |(idx, _)| idx as u32);
tokens1.push(next1);
let logits2 = decoder
.forward_one_paged(*tokens2.last().unwrap(), &encoder2, &mut paged_cache, seq2)
.unwrap();
let next2 = logits2
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
.map_or(50256, |(idx, _)| idx as u32);
tokens2.push(next2);
}
assert_eq!(tokens1.len(), 4);
assert_eq!(tokens2.len(), 4);
assert!(paged_cache.has_sequence(seq1));
assert!(paged_cache.has_sequence(seq2));
}
}
#[cfg(feature = "realizar-inference")]
mod fused_ffn_wiring_tests {
use super::*;
#[test]
fn test_decoder_block_initialize_fused_ffn() {
let d_model = 384;
let n_heads = 6;
let d_ff = 1536;
let block = DecoderBlock::new(d_model, n_heads, d_ff);
let fused = block.create_fused_ffn().expect("create fused FFN");
assert_eq!(fused.d_model, d_model);
assert_eq!(fused.d_ff, d_ff);
}
#[test]
fn test_decoder_block_fused_ffn_matches_unfused() {
let d_model = 384;
let n_heads = 6;
let d_ff = 1536;
let seq_len = 5;
let block = DecoderBlock::new(d_model, n_heads, d_ff);
let residual: Vec<f32> = (0..seq_len * d_model)
.map(|i| ((i as f32) * 0.01).sin())
.collect();
let normed = block.ln3.forward(&residual).expect("ln3");
let unfused_out = block.ffn.forward(&normed).expect("ffn");
let fused = block.create_fused_ffn().expect("create fused FFN");
let fused_out = fused.forward(&residual).expect("fused forward");
assert_eq!(unfused_out.len(), fused_out.len());
for i in 0..unfused_out.len() {
assert!(
(unfused_out[i] - fused_out[i]).abs() < 1e-5,
"Mismatch at {}: unfused {} vs fused {}",
i,
unfused_out[i],
fused_out[i]
);
}
}
#[test]
fn test_decoder_initialize_fused_ffn_all_blocks() {
let config = ModelConfig::tiny();
let mut decoder = create_decoder_with_test_weights(&config);
decoder
.initialize_fused_ffn()
.expect("initialize fused FFN should succeed");
assert_eq!(decoder.n_layers, config.n_text_layer as usize);
}
#[test]
fn test_decoder_block_forward_fused_shape() {
let d_model = 384;
let n_heads = 6;
let d_ff = 1536;
let seq_len = 5;
let enc_len = 10;
let block = DecoderBlock::new(d_model, n_heads, d_ff);
let x: Vec<f32> = (0..seq_len * d_model)
.map(|i| (i as f32 * 0.01).sin())
.collect();
let encoder_output: Vec<f32> = (0..enc_len * d_model)
.map(|i| (i as f32 * 0.02).cos())
.collect();
let unfused_out = block
.forward(&x, &encoder_output, None)
.expect("unfused forward");
let fused_out = block
.forward_fused(&x, &encoder_output, None)
.expect("fused forward");
assert_eq!(unfused_out.len(), fused_out.len());
assert_eq!(fused_out.len(), seq_len * d_model);
}
#[test]
fn test_decoder_block_forward_fused_matches_unfused() {
let d_model = 384;
let n_heads = 6;
let d_ff = 1536;
let seq_len = 5;
let enc_len = 10;
let block = DecoderBlock::new(d_model, n_heads, d_ff);
let x: Vec<f32> = (0..seq_len * d_model)
.map(|i| (i as f32 * 0.01).sin())
.collect();
let encoder_output: Vec<f32> = (0..enc_len * d_model)
.map(|i| (i as f32 * 0.02).cos())
.collect();
let unfused_out = block
.forward(&x, &encoder_output, None)
.expect("unfused forward");
let fused_out = block
.forward_fused(&x, &encoder_output, None)
.expect("fused forward");
for i in 0..unfused_out.len() {
assert!(
(unfused_out[i] - fused_out[i]).abs() < 1e-5,
"Mismatch at {}: unfused {} vs fused {}",
i,
unfused_out[i],
fused_out[i]
);
}
}
#[test]
fn test_decoder_forward_one_fused() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let mut cache = decoder.create_kv_cache();
let token = 50257_u32; let logits = decoder
.forward_one_fused(token, &encoder_output, &mut cache)
.expect("forward_one_fused");
assert_eq!(logits.len(), decoder.n_vocab);
let sum: f32 = logits.iter().map(|x: &f32| x.abs()).sum();
assert!(sum > 0.0, "Logits should not be all zeros");
}
#[test]
fn test_batched_beam_search_independence() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let batch_size = 2;
let encoder_outputs = vec![
(0..5 * d_model)
.map(|i| (i as f32).sin())
.collect::<Vec<f32>>(),
(0..5 * d_model)
.map(|i| (i as f32).cos())
.collect::<Vec<f32>>(),
];
let mut cache = BatchDecoderCache::new(
batch_size,
config.n_text_layer as usize,
d_model,
config.n_text_ctx as usize,
);
let tokens = vec![50257, 50257];
let out1 = decoder
.forward_one_batch_fused(&tokens, &encoder_outputs, &mut cache)
.unwrap();
assert_eq!(out1.logits.len(), 2);
assert_eq!(cache.get_cache(0).unwrap().seq_len(), 1);
assert_eq!(cache.get_cache(1).unwrap().seq_len(), 1);
}
#[test]
fn test_batched_beam_search_forking() {
let config = ModelConfig::tiny();
let d_model = config.n_text_state as usize;
let mut cache = BatchDecoderCache::new(
2,
config.n_text_layer as usize,
d_model,
config.n_text_ctx as usize,
);
cache.get_cache_mut(0).unwrap().self_attn_cache[0]
.append(&vec![1.0; d_model], &vec![2.0; d_model])
.unwrap();
cache.fork(0, 1).unwrap();
assert_eq!(cache.get_cache(0).unwrap().seq_len(), 1);
assert_eq!(cache.get_cache(1).unwrap().seq_len(), 1);
cache.get_cache_mut(0).unwrap().self_attn_cache[0]
.append(&vec![3.0; d_model], &vec![4.0; d_model])
.unwrap();
assert_eq!(cache.get_cache(0).unwrap().seq_len(), 2);
assert_eq!(cache.get_cache(1).unwrap().seq_len(), 1);
}
}
#[cfg(feature = "realizar-inference")]
mod speculative_tests {
use super::*;
use crate::realizar_inference::SpeculativeModel;
#[test]
fn test_speculative_decoder_wrapper_new() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let mut cache = decoder.create_kv_cache();
let wrapper = SpeculativeDecoderWrapper::new(&decoder, &encoder_output, &mut cache);
assert_eq!(wrapper.vocab_size(), config.n_vocab as usize);
assert_eq!(wrapper.eos_token(), 50257);
}
#[test]
fn test_speculative_decoder_wrapper_custom_eos() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let mut cache = decoder.create_kv_cache();
let wrapper = SpeculativeDecoderWrapper::new(&decoder, &encoder_output, &mut cache)
.with_eos_token(50256);
assert_eq!(wrapper.eos_token(), 50256);
}
#[test]
fn test_speculative_decoder_wrapper_forward() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let mut cache = decoder.create_kv_cache();
let wrapper = SpeculativeDecoderWrapper::new(&decoder, &encoder_output, &mut cache);
let tokens = [50257_u32]; let logits = wrapper.forward(&tokens).expect("forward");
assert_eq!(logits.len(), config.n_vocab as usize);
}
#[test]
fn test_speculative_decoder_wrapper_sample() {
let config = ModelConfig::tiny();
let decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let mut cache = decoder.create_kv_cache();
let wrapper = SpeculativeDecoderWrapper::new(&decoder, &encoder_output, &mut cache);
let mut logits = vec![-10.0_f32; 100];
logits[3] = 10.0;
let token_prob = wrapper.sample(&logits).expect("sample");
assert_eq!(token_prob.token, 3);
assert!(token_prob.log_prob > -1.0); }
#[test]
fn test_whisper_speculative_config_default() {
let config = WhisperSpeculativeConfig::default();
assert_eq!(config.lookahead, 4);
assert_eq!(config.max_tokens, 448);
assert_eq!(config.eos_token, 50257);
assert!(config.acceptance_threshold > 0.0);
}
#[test]
fn test_generate_speculative_short_sequence() {
let config = ModelConfig::tiny();
let draft_decoder = create_decoder_with_test_weights(&config);
let target_decoder = create_decoder_with_test_weights(&config);
let d_model = config.n_text_state as usize;
let encoder_output: Vec<f32> = (0..5 * d_model).map(|i| (i as f32).sin()).collect();
let initial_tokens = vec![50257_u32]; let spec_config = WhisperSpeculativeConfig {
lookahead: 2,
acceptance_threshold: 0.8,
max_tokens: 5,
eos_token: 50257,
};
let tokens = target_decoder
.generate_speculative(
&draft_decoder,
&encoder_output,
&initial_tokens,
&spec_config,
)
.expect("generate_speculative");
assert!(tokens.len() >= initial_tokens.len());
}
#[test]
fn test_sample_with_prob_greedy() {
let mut logits = vec![0.0_f32; 10];
logits[7] = 5.0;
let (token, prob) = sample_with_prob(&logits);
assert_eq!(token, 7);
assert!(prob > 0.9); }
#[test]
fn test_sample_with_prob_uniform() {
let logits = vec![1.0_f32; 10];
let (token, prob) = sample_with_prob(&logits);
assert!(token < 10);
assert!((prob - 0.1).abs() < 0.01);
}
}
}