#![allow(clippy::all, clippy::pedantic, clippy::restriction, clippy::nursery)]
use crate::error::{WhisperError, WhisperResult};
use crate::parallel::{parallel_map, parallel_try_map};
use crate::simd;
use trueno::Matrix;
#[cfg(feature = "realizar-inference")]
use realizar::layers::Attention as RealizarAttention;
#[cfg(feature = "realizar-inference")]
use realizar::tensor::Tensor as RealizarTensor;
#[derive(Debug, Clone)]
pub enum WeightStorage {
F32,
F16(Vec<u16>),
}
pub struct LinearWeights {
pub weight: Vec<f32>,
pub bias: Vec<f32>,
pub in_features: usize,
pub out_features: usize,
weight_transposed: Option<Vec<f32>>,
weight_matrix: Option<Matrix<f32>>,
weight_f16: Option<Vec<u16>>,
weight_i8: Option<Vec<i8>>,
weight_i8_scales: Option<Vec<f32>>,
weight_i4: Option<Vec<u8>>,
weight_i4_scales: Option<Vec<f32>>,
weight_prepacked_b: Option<trueno::blis::PrepackedB>,
}
impl Clone for LinearWeights {
fn clone(&self) -> Self {
Self {
weight: self.weight.clone(),
bias: self.bias.clone(),
in_features: self.in_features,
out_features: self.out_features,
weight_transposed: self.weight_transposed.clone(),
weight_matrix: None,
weight_f16: self.weight_f16.clone(),
weight_i8: self.weight_i8.clone(),
weight_i8_scales: self.weight_i8_scales.clone(),
weight_i4: self.weight_i4.clone(),
weight_i4_scales: self.weight_i4_scales.clone(),
weight_prepacked_b: None,
}
}
}
impl std::fmt::Debug for LinearWeights {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LinearWeights")
.field("in_features", &self.in_features)
.field("out_features", &self.out_features)
.field("weight_len", &self.weight.len())
.field("bias_len", &self.bias.len())
.field("is_finalized", &self.weight_transposed.is_some())
.field("has_matrix_cache", &self.weight_matrix.is_some())
.field("has_f16", &self.weight_f16.is_some())
.field("has_i8", &self.weight_i8.is_some())
.finish()
}
}
impl LinearWeights {
#[must_use]
pub fn new(in_features: usize, out_features: usize) -> Self {
Self {
weight: vec![0.0; out_features * in_features],
bias: vec![0.0; out_features],
in_features,
out_features,
weight_transposed: None,
weight_matrix: None,
weight_f16: None,
weight_i8: None,
weight_i8_scales: None,
weight_i4: None,
weight_i4_scales: None,
weight_prepacked_b: None,
}
}
pub fn finalize_weights(&mut self) {
if self.weight_f16.is_some() {
return;
}
let weight_t = simd::transpose(&self.weight, self.out_features, self.in_features);
const PREPACK_THRESHOLD: usize = 0;
if self.in_features * self.out_features >= PREPACK_THRESHOLD {
self.weight_prepacked_b = Some(trueno::blis::PrepackedB::pack(
&weight_t,
self.in_features,
self.out_features,
));
}
self.weight_matrix =
Matrix::from_vec(self.in_features, self.out_features, weight_t.clone()).ok();
self.weight_transposed = Some(weight_t);
}
pub fn finalize_weights_encoder(&mut self) {
let weight_t = if let Some(ref w_f16) = self.weight_f16 {
let mut buf = vec![0.0_f32; w_f16.len()];
crate::simd::dequant_f16_row(w_f16, &mut buf);
crate::simd::transpose(&buf, self.out_features, self.in_features)
} else {
crate::simd::transpose(&self.weight, self.out_features, self.in_features)
};
self.weight_prepacked_b = None;
self.weight_matrix =
trueno::Matrix::from_vec(self.in_features, self.out_features, weight_t.clone()).ok();
self.weight_transposed = Some(weight_t);
}
#[must_use]
pub fn is_finalized(&self) -> bool {
self.weight_transposed.is_some()
}
#[must_use]
pub fn weight_f16(&self) -> Option<&[u16]> {
self.weight_f16.as_deref()
}
pub fn invalidate_cache(&mut self) {
self.weight_transposed = None;
self.weight_matrix = None;
self.weight_prepacked_b = None;
}
pub fn set_weight(&mut self, values: &[f32]) {
let len = values.len().min(self.weight.len());
self.weight[..len].copy_from_slice(&values[..len]);
self.invalidate_cache();
}
pub fn set_bias(&mut self, values: &[f32]) {
let len = values.len().min(self.bias.len());
self.bias[..len].copy_from_slice(&values[..len]);
}
pub fn set_weight_f16(&mut self, values: &[u16]) {
self.weight_f16 = Some(values.to_vec());
self.weight = Vec::new();
self.invalidate_cache();
}
pub fn convert_to_f16(&mut self) {
if self.weight_f16.is_some() || self.weight.is_empty() {
return;
}
self.weight_f16 = Some(simd::quant_f32_to_f16(&self.weight));
self.weight = Vec::new();
self.invalidate_cache();
}
pub fn convert_to_i8(&mut self) {
if self.weight_i8.is_some() {
return;
}
let rows = self.out_features;
let cols = self.in_features;
if let Some(ref w_f16) = self.weight_f16 {
let mut all_i8 = Vec::with_capacity(rows * cols);
let mut scales = Vec::with_capacity(rows);
let mut row_buf = vec![0.0_f32; cols];
for r in 0..rows {
let offset = r * cols;
for (j, v) in row_buf.iter_mut().enumerate() {
*v = half::f16::from_bits(w_f16[offset + j]).to_f32();
}
let (q_row, scale) = simd::quant_f32_row_to_i8(&row_buf);
all_i8.extend_from_slice(&q_row);
scales.push(scale);
}
self.weight_i8 = Some(all_i8);
self.weight_i8_scales = Some(scales);
} else if !self.weight.is_empty() {
let mut all_i8 = Vec::with_capacity(rows * cols);
let mut scales = Vec::with_capacity(rows);
for r in 0..rows {
let offset = r * cols;
let (q_row, scale) = simd::quant_f32_row_to_i8(&self.weight[offset..offset + cols]);
all_i8.extend_from_slice(&q_row);
scales.push(scale);
}
self.weight_i8 = Some(all_i8);
self.weight_i8_scales = Some(scales);
}
}
pub fn convert_to_i4(&mut self) {
if self.weight_i4.is_some() {
return;
}
let rows = self.out_features;
let cols = self.in_features;
let group_size = 128;
if let Some(ref w_f16) = self.weight_f16 {
let mut all_i4 = Vec::with_capacity(rows * cols / 2);
let mut all_scales = Vec::with_capacity(rows * (cols / group_size));
let mut row_buf = vec![0.0_f32; cols];
for r in 0..rows {
let offset = r * cols;
for (j, v) in row_buf.iter_mut().enumerate() {
*v = half::f16::from_bits(w_f16[offset + j]).to_f32();
}
let (q_row, scales) = simd::quant_f32_row_to_i4(&row_buf, group_size);
all_i4.extend_from_slice(&q_row);
all_scales.extend_from_slice(&scales);
}
self.weight_i4 = Some(all_i4);
self.weight_i4_scales = Some(all_scales);
} else if !self.weight.is_empty() {
let mut all_i4 = Vec::with_capacity(rows * cols / 2);
let mut all_scales = Vec::with_capacity(rows * (cols / group_size));
for r in 0..rows {
let offset = r * cols;
let (q_row, scales) = simd::quant_f32_row_to_i4(&self.weight[offset..offset + cols], group_size);
all_i4.extend_from_slice(&q_row);
all_scales.extend_from_slice(&scales);
}
self.weight_i4 = Some(all_i4);
self.weight_i4_scales = Some(all_scales);
}
}
#[must_use]
pub fn is_f16(&self) -> bool {
self.weight_f16.is_some()
}
#[must_use]
pub fn is_i8(&self) -> bool {
self.weight_i8.is_some()
}
#[must_use]
pub fn storage_type(&self) -> WeightStorage {
if let Some(ref f16) = self.weight_f16 {
WeightStorage::F16(f16.clone())
} else {
WeightStorage::F32
}
}
pub fn forward(&self, input: &[f32], seq_len: usize) -> WhisperResult<Vec<f32>> {
if input.len() % (seq_len * self.in_features) != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
let batch_size = input.len() / (seq_len * self.in_features);
debug_assert!(batch_size > 0, "batch_size must be positive");
let mut output = vec![0.0_f32; batch_size * seq_len * self.out_features];
for b in 0..batch_size {
for s in 0..seq_len {
for o in 0..self.out_features {
let mut sum = self.bias[o];
for i in 0..self.in_features {
let input_idx = b * seq_len * self.in_features + s * self.in_features + i;
let weight_idx = o * self.in_features + i;
sum += input[input_idx] * self.weight[weight_idx];
}
let output_idx = b * seq_len * self.out_features + s * self.out_features + o;
output[output_idx] = sum;
}
}
}
debug_assert_eq!(
output.len(),
batch_size * seq_len * self.out_features,
"output dimensions must match batch × seq × out_features"
);
Ok(output)
}
pub fn forward_simd(&self, input: &[f32], seq_len: usize) -> WhisperResult<Vec<f32>> {
if input.len() % (seq_len * self.in_features) != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
let batch_size = input.len() / (seq_len * self.in_features);
let total_tokens = batch_size * seq_len;
if let Some(ref w_f16) = self.weight_f16 {
const FP16_MATVEC_BATCH_LIMIT: usize = 8;
let mut output = if total_tokens <= FP16_MATVEC_BATCH_LIMIT {
let mut out = vec![0.0_f32; total_tokens * self.out_features];
for t in 0..total_tokens {
let tok_in = &input[t * self.in_features..(t + 1) * self.in_features];
let tok_out = &mut out[t * self.out_features..(t + 1) * self.out_features];
simd::tiled_matvec_f16_into(
w_f16,
tok_in,
tok_out,
self.out_features,
self.in_features,
);
}
out
} else if let Some(ref prepacked) = self.weight_prepacked_b {
simd::matmul_with_prepacked(
input,
prepacked,
total_tokens,
self.in_features,
self.out_features,
)
} else if let Some(weight_matrix) = &self.weight_matrix {
simd::matmul_with_matrix(input, weight_matrix, total_tokens, self.in_features)
} else if let Some(weight_t) = &self.weight_transposed {
simd::matmul(
input,
weight_t,
total_tokens,
self.in_features,
self.out_features,
)
} else {
let mut buf = vec![0.0_f32; w_f16.len()];
simd::dequant_f16_row(w_f16, &mut buf);
let weight_t = simd::transpose(&buf, self.out_features, self.in_features);
simd::matmul(
input,
&weight_t,
total_tokens,
self.in_features,
self.out_features,
)
};
simd::broadcast_add_inplace(&mut output, &self.bias, total_tokens, self.out_features);
return Ok(output);
}
let mut output = if total_tokens == 1 {
simd::tiled_matvec(&self.weight, input, self.out_features, self.in_features)
} else if let Some(ref prepacked) = self.weight_prepacked_b {
simd::matmul_with_prepacked(
input,
prepacked,
total_tokens,
self.in_features,
self.out_features,
)
} else if let Some(prepacked_b) = &self.weight_prepacked_b {
simd::matmul_with_prepacked(
input,
prepacked_b,
total_tokens,
self.in_features,
self.out_features,
)
} else {
let mut output = vec![0.0_f32; total_tokens * self.out_features];
crate::simd::optimized::tiled_matmul_into(
&self.weight,
input,
&mut output,
total_tokens,
self.out_features,
self.in_features,
);
output
};
simd::broadcast_add_inplace(&mut output, &self.bias, total_tokens, self.out_features);
Ok(output)
}
pub fn forward_simd_into(
&self,
input: &[f32],
seq_len: usize,
output: &mut [f32],
) -> WhisperResult<()> {
if input.len() % (seq_len * self.in_features) != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
let batch_size = input.len() / (seq_len * self.in_features);
let total_tokens = batch_size * seq_len;
if let (Some(ref w_i4), Some(ref scales)) = (&self.weight_i4, &self.weight_i4_scales) {
let group_size = 128;
for t in 0..total_tokens {
let tok_in = &input[t * self.in_features..(t + 1) * self.in_features];
let tok_out = &mut output[t * self.out_features..(t + 1) * self.out_features];
simd::tiled_matvec_i4_into(
w_i4,
scales,
tok_in,
tok_out,
self.out_features,
self.in_features,
group_size,
);
}
simd::broadcast_add_inplace(output, &self.bias, total_tokens, self.out_features);
return Ok(());
}
if let (Some(ref w_i8), Some(ref scales)) = (&self.weight_i8, &self.weight_i8_scales) {
for t in 0..total_tokens {
let tok_in = &input[t * self.in_features..(t + 1) * self.in_features];
let tok_out = &mut output[t * self.out_features..(t + 1) * self.out_features];
simd::tiled_matvec_i8_into(
w_i8,
scales,
tok_in,
tok_out,
self.out_features,
self.in_features,
);
}
simd::broadcast_add_inplace(output, &self.bias, total_tokens, self.out_features);
return Ok(());
}
if let Some(ref w_f16) = self.weight_f16 {
const FP16_MATVEC_BATCH_LIMIT: usize = 8;
if total_tokens <= FP16_MATVEC_BATCH_LIMIT {
for t in 0..total_tokens {
let tok_in = &input[t * self.in_features..(t + 1) * self.in_features];
let tok_out = &mut output[t * self.out_features..(t + 1) * self.out_features];
simd::tiled_matvec_f16_into(
w_f16,
tok_in,
tok_out,
self.out_features,
self.in_features,
);
}
} else if let Some(ref prepacked) = self.weight_prepacked_b {
let tmp = simd::matmul_with_prepacked(
input,
prepacked,
total_tokens,
self.in_features,
self.out_features,
);
output[..tmp.len()].copy_from_slice(&tmp);
} else if let Some(weight_matrix) = &self.weight_matrix {
let tmp = simd::matmul_with_matrix(input, weight_matrix, total_tokens, self.in_features);
output[..tmp.len()].copy_from_slice(&tmp);
} else if let Some(weight_t) = &self.weight_transposed {
let tmp = simd::matmul(
input,
weight_t,
total_tokens,
self.in_features,
self.out_features,
);
output[..tmp.len()].copy_from_slice(&tmp);
} else {
let mut buf = vec![0.0_f32; w_f16.len()];
simd::dequant_f16_row(w_f16, &mut buf);
let weight_t = simd::transpose(&buf, self.out_features, self.in_features);
let tmp = simd::matmul(
input,
&weight_t,
total_tokens,
self.in_features,
self.out_features,
);
output[..tmp.len()].copy_from_slice(&tmp);
}
simd::broadcast_add_inplace(output, &self.bias, total_tokens, self.out_features);
return Ok(());
}
if total_tokens == 1 {
simd::tiled_matvec_into(
&self.weight,
input,
output,
self.out_features,
self.in_features,
);
} else if let Some(ref prepacked) = self.weight_prepacked_b {
let tmp = simd::matmul_with_prepacked(
input,
prepacked,
total_tokens,
self.in_features,
self.out_features,
);
output[..tmp.len()].copy_from_slice(&tmp);
} else if let Some(prepacked_b) = &self.weight_prepacked_b {
let tmp = simd::matmul_with_prepacked(
input,
prepacked_b,
total_tokens,
self.in_features,
self.out_features,
);
output[..tmp.len()].copy_from_slice(&tmp);
} else {
crate::simd::optimized::tiled_matmul_into(
&self.weight,
input,
output,
total_tokens,
self.out_features,
self.in_features,
);
}
simd::broadcast_add_inplace(output, &self.bias, total_tokens, self.out_features);
Ok(())
}
}
pub const FLASH_ATTENTION_BLOCK_SIZE: usize = 128;
pub const FLASH_ATTENTION_THRESHOLD: usize = 128;
#[derive(Debug, Clone, Copy)]
pub struct FlashAttentionConfig {
pub seq_len: usize,
pub kv_len: usize,
pub d_head: usize,
pub block_size: usize,
}
impl FlashAttentionConfig {
#[must_use]
pub const fn new(seq_len: usize, kv_len: usize, d_head: usize, block_size: usize) -> Self {
Self {
seq_len,
kv_len,
d_head,
block_size,
}
}
#[must_use]
#[allow(dead_code)] pub const fn with_default_block_size(seq_len: usize, kv_len: usize, d_head: usize) -> Self {
Self::new(seq_len, kv_len, d_head, FLASH_ATTENTION_BLOCK_SIZE)
}
}
struct BlockContext {
q_idx: usize,
kv_block_start: usize,
kv_block_end: usize,
scale: f32,
}
#[inline]
fn compute_block_scores(
query: &[f32],
key: &[f32],
config: &FlashAttentionConfig,
ctx: &BlockContext,
mask: Option<&[f32]>,
) -> Vec<f32> {
let mut scores = Vec::with_capacity(ctx.kv_block_end - ctx.kv_block_start);
for k_idx in ctx.kv_block_start..ctx.kv_block_end {
let mut dot = 0.0_f32;
for d in 0..config.d_head {
dot += query[ctx.q_idx * config.d_head + d] * key[k_idx * config.d_head + d];
}
let mut score = dot * ctx.scale;
if let Some(m) = mask {
score += m[ctx.q_idx * config.kv_len + k_idx];
}
scores.push(score);
}
scores
}
#[inline]
fn update_output_with_block(
output: &mut [f32],
row_sum: &mut f32,
row_max: &mut f32,
block_scores: &[f32],
value: &[f32],
ctx: &BlockContext,
d_head: usize,
) {
let block_max = block_scores
.iter()
.fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let prev_max = *row_max;
let new_max = prev_max.max(block_max);
let scale_prev = (prev_max - new_max).exp();
*row_sum *= scale_prev;
let out_start = ctx.q_idx * d_head;
for d in 0..d_head {
output[out_start + d] *= scale_prev;
}
for (local_k_idx, &score) in block_scores.iter().enumerate() {
let k_idx = ctx.kv_block_start + local_k_idx;
let exp_score = (score - new_max).exp();
*row_sum += exp_score;
for d in 0..d_head {
output[out_start + d] += exp_score * value[k_idx * d_head + d];
}
}
*row_max = new_max;
}
#[inline]
fn normalize_row(output: &mut [f32], row_sum: f32, q_idx: usize, d_head: usize) {
let inv_sum = if row_sum > 1e-10 { 1.0 / row_sum } else { 0.0 };
let start = q_idx * d_head;
for d in 0..d_head {
output[start + d] *= inv_sum;
}
}
#[must_use]
pub fn flash_attention(
query: &[f32],
key: &[f32],
value: &[f32],
config: FlashAttentionConfig,
mask: Option<&[f32]>,
) -> Vec<f32> {
let scale = 1.0 / (config.d_head as f32).sqrt();
let mut output = vec![0.0_f32; config.seq_len * config.d_head];
let mut row_max = vec![f32::NEG_INFINITY; config.seq_len];
let mut row_sum = vec![0.0_f32; config.seq_len];
for kv_block_start in (0..config.kv_len).step_by(config.block_size) {
let kv_block_end = (kv_block_start + config.block_size).min(config.kv_len);
for q_idx in 0..config.seq_len {
let ctx = BlockContext {
q_idx,
kv_block_start,
kv_block_end,
scale,
};
let block_scores = compute_block_scores(query, key, &config, &ctx, mask);
update_output_with_block(
&mut output,
&mut row_sum[q_idx],
&mut row_max[q_idx],
&block_scores,
value,
&ctx,
config.d_head,
);
}
}
for (q_idx, &sum) in row_sum.iter().enumerate() {
normalize_row(&mut output, sum, q_idx, config.d_head);
}
output
}
#[inline]
fn compute_block_scores_simd(
query: &[f32],
key: &[f32],
config: &FlashAttentionConfig,
ctx: &BlockContext,
mask: Option<&[f32]>,
) -> Vec<f32> {
let q_offset = ctx.q_idx * config.d_head;
let mut scores = Vec::with_capacity(ctx.kv_block_end - ctx.kv_block_start);
for k_idx in ctx.kv_block_start..ctx.kv_block_end {
let k_offset = k_idx * config.d_head;
let dot = simd::dot(
&query[q_offset..q_offset + config.d_head],
&key[k_offset..k_offset + config.d_head],
);
let mut score = dot * ctx.scale;
if let Some(m) = mask {
score += m[ctx.q_idx * config.kv_len + k_idx];
}
scores.push(score);
}
scores
}
#[inline]
fn update_output_simd(
output: &mut [f32],
row_sum: &mut f32,
row_max: &mut f32,
block_scores: &[f32],
value: &[f32],
ctx: &BlockContext,
d_head: usize,
) {
let block_max = simd::max_element(block_scores);
let prev_max = *row_max;
let new_max = prev_max.max(block_max);
let scale_prev = (prev_max - new_max).exp();
*row_sum *= scale_prev;
let q_offset = ctx.q_idx * d_head;
simd::scale_inplace(&mut output[q_offset..q_offset + d_head], scale_prev);
for (local_k_idx, &score) in block_scores.iter().enumerate() {
let k_idx = ctx.kv_block_start + local_k_idx;
let exp_score = (score - new_max).exp();
*row_sum += exp_score;
let v_offset = k_idx * d_head;
simd::axpy(
exp_score,
&value[v_offset..v_offset + d_head],
&mut output[q_offset..q_offset + d_head],
);
}
*row_max = new_max;
}
#[must_use]
pub fn flash_attention_simd(
query: &[f32],
key: &[f32],
value: &[f32],
config: FlashAttentionConfig,
mask: Option<&[f32]>,
) -> Vec<f32> {
let scale = 1.0 / (config.d_head as f32).sqrt();
let mut output = vec![0.0_f32; config.seq_len * config.d_head];
let mut row_max = vec![f32::NEG_INFINITY; config.seq_len];
let mut row_sum = vec![0.0_f32; config.seq_len];
for kv_block_start in (0..config.kv_len).step_by(config.block_size) {
let kv_block_end = (kv_block_start + config.block_size).min(config.kv_len);
for q_idx in 0..config.seq_len {
let ctx = BlockContext {
q_idx,
kv_block_start,
kv_block_end,
scale,
};
let block_scores = compute_block_scores_simd(query, key, &config, &ctx, mask);
update_output_simd(
&mut output,
&mut row_sum[q_idx],
&mut row_max[q_idx],
&block_scores,
value,
&ctx,
config.d_head,
);
}
}
for (q_idx, &sum) in row_sum.iter().enumerate() {
let inv_sum = if sum > 1e-10 { 1.0 / sum } else { 0.0 };
let start = q_idx * config.d_head;
let end = (q_idx + 1) * config.d_head;
simd::scale_inplace(&mut output[start..end], inv_sum);
}
output
}
#[cfg(feature = "parallel")]
#[must_use]
pub fn flash_attention_simd_parallel(
query: &[f32],
key: &[f32],
value: &[f32],
config: FlashAttentionConfig,
mask: Option<&[f32]>,
) -> Vec<f32> {
use rayon::prelude::*;
let scale = 1.0 / (config.d_head as f32).sqrt();
let d_head = config.d_head;
let kv_len = config.kv_len;
let block_size = config.block_size;
let mut output = vec![0.0_f32; config.seq_len * d_head];
output
.par_chunks_mut(d_head)
.enumerate()
.for_each(|(q_idx, out_row)| {
let mut row_max = f32::NEG_INFINITY;
let mut row_sum = 0.0_f32;
let mut block_scores = Vec::with_capacity(block_size);
let q_offset = q_idx * d_head;
for kv_block_start in (0..kv_len).step_by(block_size) {
let kv_block_end = (kv_block_start + block_size).min(kv_len);
block_scores.clear();
for k_idx in kv_block_start..kv_block_end {
let k_offset = k_idx * d_head;
let dot = simd::dot_nalloc(
&query[q_offset..q_offset + d_head],
&key[k_offset..k_offset + d_head],
);
let mut score = dot * scale;
if let Some(m) = mask {
score += m[q_idx * kv_len + k_idx];
}
block_scores.push(score);
}
let block_max = block_scores
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let new_max = row_max.max(block_max);
let scale_prev = (row_max - new_max).exp();
row_sum *= scale_prev;
simd::scale_inplace(out_row, scale_prev);
for (local_k_idx, &score) in block_scores.iter().enumerate() {
let k_idx = kv_block_start + local_k_idx;
let exp_score = (score - new_max).exp();
row_sum += exp_score;
let v_offset = k_idx * d_head;
simd::axpy(exp_score, &value[v_offset..v_offset + d_head], out_row);
}
row_max = new_max;
}
let inv_sum = if row_sum > 1e-10 { 1.0 / row_sum } else { 0.0 };
simd::scale_inplace(out_row, inv_sum);
});
output
}
#[derive(Debug, Clone)]
pub struct MultiHeadAttention {
n_heads: usize,
d_model: usize,
d_head: usize,
w_q: LinearWeights,
w_k: LinearWeights,
w_v: LinearWeights,
w_o: LinearWeights,
scale: f32,
w_qkv_f16: Option<Vec<u16>>,
b_qkv: Option<Vec<f32>>,
}
impl MultiHeadAttention {
#[must_use]
pub fn new(n_heads: usize, d_model: usize) -> Self {
assert!(
d_model % n_heads == 0,
"d_model ({d_model}) must be divisible by n_heads ({n_heads})"
);
let d_head = d_model / n_heads;
Self {
n_heads,
d_model,
d_head,
w_q: LinearWeights::new(d_model, d_model),
w_k: LinearWeights::new(d_model, d_model),
w_v: LinearWeights::new(d_model, d_model),
w_o: LinearWeights::new(d_model, d_model),
scale: 1.0 / (d_head as f32).sqrt(),
w_qkv_f16: None,
b_qkv: None,
}
}
pub fn scaled_dot_product_attention(
&self,
query: &[f32],
key: &[f32],
value: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = query.len() / self.d_head;
let kv_len = key.len() / self.d_head;
if query.len() % self.d_head != 0 {
return Err(WhisperError::Model("query size mismatch".into()));
}
if key.len() % self.d_head != 0 || value.len() % self.d_head != 0 {
return Err(WhisperError::Model("key/value size mismatch".into()));
}
if key.len() != value.len() {
return Err(WhisperError::Model(
"key and value must have same length".into(),
));
}
let mut scores = vec![0.0_f32; seq_len * kv_len];
for q_idx in 0..seq_len {
for k_idx in 0..kv_len {
let mut dot = 0.0_f32;
for d in 0..self.d_head {
dot += query[q_idx * self.d_head + d] * key[k_idx * self.d_head + d];
}
scores[q_idx * kv_len + k_idx] = dot * self.scale;
}
}
if let Some(m) = mask {
if m.len() != seq_len * kv_len {
return Err(WhisperError::Model("mask size mismatch".into()));
}
for i in 0..scores.len() {
scores[i] += m[i];
}
}
Self::apply_row_softmax(&mut scores, seq_len, kv_len);
let mut output = vec![0.0_f32; seq_len * self.d_head];
for q_idx in 0..seq_len {
for d in 0..self.d_head {
let mut sum = 0.0_f32;
for k_idx in 0..kv_len {
sum += scores[q_idx * kv_len + k_idx] * value[k_idx * self.d_head + d];
}
output[q_idx * self.d_head + d] = sum;
}
}
Ok(output)
}
fn apply_row_softmax(scores: &mut [f32], seq_len: usize, kv_len: usize) {
for q_idx in 0..seq_len {
let row_start = q_idx * kv_len;
let row_end = row_start + kv_len;
let max_score = scores[row_start..row_end]
.iter()
.fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut sum = 0.0_f32;
for k_idx in 0..kv_len {
let exp_val = (scores[row_start + k_idx] - max_score).exp();
scores[row_start + k_idx] = exp_val;
sum += exp_val;
}
let inv_sum = if sum > 1e-10 { 1.0 / sum } else { 0.0 };
for k_idx in 0..kv_len {
scores[row_start + k_idx] *= inv_sum;
}
}
}
pub fn scaled_dot_product_attention_simd(
&self,
query: &[f32],
key: &[f32],
value: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = query.len() / self.d_head;
if query.len() % self.d_head != 0 {
return Err(WhisperError::Model("query size mismatch".into()));
}
if key.len() % self.d_head != 0 || value.len() % self.d_head != 0 {
return Err(WhisperError::Model("key/value size mismatch".into()));
}
if key.len() != value.len() {
return Err(WhisperError::Model(
"key and value must have same length".into(),
));
}
let output =
simd::scaled_dot_product_attention(query, key, value, seq_len, self.d_head, mask);
Ok(output)
}
#[must_use]
pub fn causal_mask(seq_len: usize) -> Vec<f32> {
let mut mask = vec![0.0_f32; seq_len * seq_len];
for i in 0..seq_len {
for j in 0..seq_len {
if j > i {
mask[i * seq_len + j] = f32::NEG_INFINITY;
}
}
}
mask
}
pub fn forward(&self, x: &[f32], mask: Option<&[f32]>) -> WhisperResult<Vec<f32>> {
self.forward_cross_dispatch(x, x, mask)
}
#[cfg(feature = "realizar-inference")]
pub fn forward_cross_dispatch(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
self.forward_cross_optimal(x, context, mask)
}
#[cfg(not(feature = "realizar-inference"))]
pub fn forward_cross_dispatch(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if seq_len > FLASH_ATTENTION_THRESHOLD || kv_len > FLASH_ATTENTION_THRESHOLD {
self.forward_cross_flash(x, context, mask, FLASH_ATTENTION_BLOCK_SIZE)
} else if cfg!(feature = "simd") {
self.forward_cross_simd(x, context, mask)
} else {
self.forward_cross(x, context, mask)
}
}
pub fn forward_cross(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
if cfg!(feature = "simd") {
self.forward_cross_simd(x, context, mask)
} else {
self.forward_cross_scalar(x, context, mask)
}
}
fn forward_cross_scalar(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if x.len() % self.d_model != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
if context.len() % self.d_model != 0 {
return Err(WhisperError::Model("context size mismatch".into()));
}
let q = self.w_q.forward(x, seq_len)?;
let k = self.w_k.forward(context, kv_len)?;
let v = self.w_v.forward(context, kv_len)?;
let head_outputs = parallel_try_map(0..self.n_heads, |head| {
let q_head = self.extract_head(&q, seq_len, head);
let k_head = self.extract_head(&k, kv_len, head);
let v_head = self.extract_head(&v, kv_len, head);
self.scaled_dot_product_attention(&q_head, &k_head, &v_head, mask)
})?;
let concat = self.concat_heads(&head_outputs, seq_len);
self.w_o.forward(&concat, seq_len)
}
fn forward_cross_simd(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if x.len() % self.d_model != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
if context.len() % self.d_model != 0 {
return Err(WhisperError::Model("context size mismatch".into()));
}
#[cfg(feature = "parallel")]
let (q, k, v) = {
let (q, (k, v)) = rayon::join(
|| self.w_q.forward_simd(x, seq_len),
|| {
rayon::join(
|| self.w_k.forward_simd(context, kv_len),
|| self.w_v.forward_simd(context, kv_len),
)
},
);
(q?, k?, v?)
};
#[cfg(not(feature = "parallel"))]
let (q, k, v) = {
let q = self.w_q.forward_simd(x, seq_len)?;
let k = self.w_k.forward_simd(context, kv_len)?;
let v = self.w_v.forward_simd(context, kv_len)?;
(q, k, v)
};
let head_outputs = parallel_try_map(0..self.n_heads, |head| {
let q_head = self.extract_head(&q, seq_len, head);
let k_head = self.extract_head(&k, kv_len, head);
let v_head = self.extract_head(&v, kv_len, head);
self.scaled_dot_product_attention_simd(&q_head, &k_head, &v_head, mask)
})?;
let concat = self.concat_heads(&head_outputs, seq_len);
self.w_o.forward_simd(&concat, seq_len)
}
pub fn forward_cross_flash(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
block_size: usize,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if x.len() % self.d_model != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
if context.len() % self.d_model != 0 {
return Err(WhisperError::Model("context size mismatch".into()));
}
let q = self.w_q.forward_simd(x, seq_len)?;
let k = self.w_k.forward_simd(context, kv_len)?;
let v = self.w_v.forward_simd(context, kv_len)?;
#[cfg(feature = "parallel")]
let head_outputs: Vec<Vec<f32>> = (0..self.n_heads)
.map(|head| {
let q_head = self.extract_head(&q, seq_len, head);
let k_head = self.extract_head(&k, kv_len, head);
let v_head = self.extract_head(&v, kv_len, head);
let config = FlashAttentionConfig::new(seq_len, kv_len, self.d_head, block_size);
flash_attention_simd_parallel(&q_head, &k_head, &v_head, config, mask)
})
.collect();
#[cfg(not(feature = "parallel"))]
let head_outputs: Vec<Vec<f32>> = (0..self.n_heads)
.map(|head| {
let q_head = self.extract_head(&q, seq_len, head);
let k_head = self.extract_head(&k, kv_len, head);
let v_head = self.extract_head(&v, kv_len, head);
let config = FlashAttentionConfig::new(seq_len, kv_len, self.d_head, block_size);
if cfg!(feature = "simd") {
flash_attention_simd(&q_head, &k_head, &v_head, config, mask)
} else {
flash_attention(&q_head, &k_head, &v_head, config, mask)
}
})
.collect();
let concat = self.concat_heads(&head_outputs, seq_len);
self.w_o.forward_simd(&concat, seq_len)
}
pub fn forward_cross_auto(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if seq_len > FLASH_ATTENTION_THRESHOLD || kv_len > FLASH_ATTENTION_THRESHOLD {
self.forward_cross_flash(x, context, mask, FLASH_ATTENTION_BLOCK_SIZE)
} else {
self.forward_cross(x, context, mask)
}
}
#[cfg(feature = "realizar-inference")]
pub fn forward_cross_flash_v2(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if x.len() % self.d_model != 0 {
return Err(WhisperError::Model("input size mismatch".into()));
}
if context.len() % self.d_model != 0 {
return Err(WhisperError::Model("context size mismatch".into()));
}
let q = self.w_q.forward_simd(x, seq_len)?;
let k = self.w_k.forward_simd(context, kv_len)?;
let v = self.w_v.forward_simd(context, kv_len)?;
let head_outputs = parallel_map(0..self.n_heads, |head| {
let q_head = self.extract_head(&q, seq_len, head);
let k_head = self.extract_head(&k, kv_len, head);
let v_head = self.extract_head(&v, kv_len, head);
let q_tensor = RealizarTensor::from_vec(vec![seq_len, self.d_head], q_head)
.expect("valid Q tensor");
let k_tensor = RealizarTensor::from_vec(vec![kv_len, self.d_head], k_head)
.expect("valid K tensor");
let v_tensor = RealizarTensor::from_vec(vec![kv_len, self.d_head], v_head)
.expect("valid V tensor");
let attn = RealizarAttention::new(self.d_head).expect("valid Attention");
let result = attn
.flash_forward_v2(&q_tensor, &k_tensor, &v_tensor, FLASH_ATTENTION_BLOCK_SIZE)
.expect("FlashAttention-2 forward");
let mut output = result.data().to_vec();
if let Some(m) = mask {
for (i, out_val) in output.iter_mut().enumerate() {
let q_idx = i / self.d_head;
if q_idx < seq_len {
let mask_row_start = q_idx * kv_len;
let all_masked = (0..kv_len).all(|k| m[mask_row_start + k] < -1e9);
if all_masked {
*out_val = 0.0;
}
}
}
}
output
});
let concat = self.concat_heads(&head_outputs, seq_len);
self.w_o.forward_simd(&concat, seq_len)
}
#[cfg(feature = "realizar-inference")]
pub fn forward_cross_optimal(
&self,
x: &[f32],
context: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<Vec<f32>> {
let seq_len = x.len() / self.d_model;
let kv_len = context.len() / self.d_model;
if seq_len > FLASH_ATTENTION_THRESHOLD || kv_len > FLASH_ATTENTION_THRESHOLD {
self.forward_cross_flash(x, context, mask, FLASH_ATTENTION_BLOCK_SIZE)
} else if cfg!(feature = "simd") {
self.forward_cross_simd(x, context, mask)
} else {
self.forward_cross(x, context, mask)
}
}
pub fn forward_streaming(
&self,
x: &[f32],
cached_key: &[f32],
cached_value: &[f32],
mask: Option<&[f32]>,
) -> WhisperResult<(Vec<f32>, Vec<f32>, Vec<f32>)> {
let seq_len = x.len() / self.d_model;
let cache_len = if cached_key.is_empty() {
0
} else {
cached_key.len() / self.d_model
};
let q = self.w_q.forward(x, seq_len)?;
let new_k = self.w_k.forward(x, seq_len)?;
let new_v = self.w_v.forward(x, seq_len)?;
let head_outputs = parallel_try_map(0..self.n_heads, |head| {
let head_q = self.extract_head(&q, seq_len, head);
let (head_k, head_v, total_kv_len) = if cache_len > 0 {
let cached_head_k = self.extract_head(cached_key, cache_len, head);
let cached_head_v = self.extract_head(cached_value, cache_len, head);
let new_head_k = self.extract_head(&new_k, seq_len, head);
let new_head_v = self.extract_head(&new_v, seq_len, head);
let mut combined_k = cached_head_k;
combined_k.extend_from_slice(&new_head_k);
let mut combined_v = cached_head_v;
combined_v.extend_from_slice(&new_head_v);
(combined_k, combined_v, cache_len + seq_len)
} else {
let new_head_k = self.extract_head(&new_k, seq_len, head);
let new_head_v = self.extract_head(&new_v, seq_len, head);
(new_head_k, new_head_v, seq_len)
};
if total_kv_len > FLASH_ATTENTION_THRESHOLD {
let config = FlashAttentionConfig::new(
seq_len,
total_kv_len,
self.d_head,
FLASH_ATTENTION_BLOCK_SIZE,
);
Ok(flash_attention_simd(
&head_q, &head_k, &head_v, config, mask,
))
} else {
self.scaled_dot_product_attention(&head_q, &head_k, &head_v, mask)
}
})?;
let concat = self.concat_heads(&head_outputs, seq_len);
let output = self.w_o.forward(&concat, seq_len)?;
Ok((output, new_k, new_v))
}
#[allow(dead_code)] pub fn forward_self_streaming(
&self,
x: &[f32],
cached_key: &[f32],
cached_value: &[f32],
) -> WhisperResult<(Vec<f32>, Vec<f32>, Vec<f32>)> {
let seq_len = x.len() / self.d_model;
let cache_len = cached_key.len() / self.d_model;
let mask = if seq_len > 1 {
let total_len = cache_len + seq_len;
let mut mask_data = vec![0.0_f32; seq_len * total_len];
for q in 0..seq_len {
let max_attend = cache_len + q + 1;
for k in max_attend..total_len {
mask_data[q * total_len + k] = f32::NEG_INFINITY;
}
}
Some(mask_data)
} else {
None
};
self.forward_streaming(x, cached_key, cached_value, mask.as_deref())
}
fn extract_head(&self, tensor: &[f32], seq_len: usize, head: usize) -> Vec<f32> {
let mut head_data = vec![0.0_f32; seq_len * self.d_head];
for s in 0..seq_len {
let src_offset = s * self.d_model + head * self.d_head;
let dst_offset = s * self.d_head;
head_data[dst_offset..dst_offset + self.d_head]
.copy_from_slice(&tensor[src_offset..src_offset + self.d_head]);
}
head_data
}
fn concat_heads(&self, heads: &[Vec<f32>], seq_len: usize) -> Vec<f32> {
let mut concat = vec![0.0_f32; seq_len * self.d_model];
for (head, head_data) in heads.iter().enumerate() {
for s in 0..seq_len {
let src_offset = s * self.d_head;
let dst_offset = s * self.d_model + head * self.d_head;
concat[dst_offset..dst_offset + self.d_head]
.copy_from_slice(&head_data[src_offset..src_offset + self.d_head]);
}
}
concat
}
#[must_use]
pub const fn n_heads(&self) -> usize {
self.n_heads
}
#[must_use]
pub const fn d_model(&self) -> usize {
self.d_model
}
#[must_use]
pub const fn d_head(&self) -> usize {
self.d_head
}
#[must_use]
pub fn scale(&self) -> f32 {
self.scale
}
#[must_use]
pub const fn w_q(&self) -> &LinearWeights {
&self.w_q
}
#[must_use]
pub const fn w_k(&self) -> &LinearWeights {
&self.w_k
}
#[must_use]
pub const fn w_v(&self) -> &LinearWeights {
&self.w_v
}
#[must_use]
pub const fn w_o(&self) -> &LinearWeights {
&self.w_o
}
pub fn set_query_weight(&mut self, values: &[f32]) {
self.w_q.set_weight(values);
}
pub fn set_key_weight(&mut self, values: &[f32]) {
self.w_k.set_weight(values);
}
pub fn set_value_weight(&mut self, values: &[f32]) {
self.w_v.set_weight(values);
}
pub fn set_out_weight(&mut self, values: &[f32]) {
self.w_o.set_weight(values);
}
pub fn set_query_bias(&mut self, values: &[f32]) {
self.w_q.set_bias(values);
}
pub fn set_key_bias(&mut self, values: &[f32]) {
self.w_k.set_bias(values);
}
pub fn set_value_bias(&mut self, values: &[f32]) {
self.w_v.set_bias(values);
}
pub fn set_out_bias(&mut self, values: &[f32]) {
self.w_o.set_bias(values);
}
pub fn set_query_weight_f16(&mut self, values: &[u16]) {
self.w_q.set_weight_f16(values);
}
pub fn set_key_weight_f16(&mut self, values: &[u16]) {
self.w_k.set_weight_f16(values);
}
pub fn set_value_weight_f16(&mut self, values: &[u16]) {
self.w_v.set_weight_f16(values);
}
pub fn set_out_weight_f16(&mut self, values: &[u16]) {
self.w_o.set_weight_f16(values);
}
pub fn convert_to_f16(&mut self) {
self.w_q.convert_to_f16();
self.w_k.convert_to_f16();
self.w_v.convert_to_f16();
self.w_o.convert_to_f16();
}
pub fn convert_to_i8(&mut self) {
self.w_q.convert_to_i8();
self.w_k.convert_to_i8();
self.w_v.convert_to_i8();
self.w_o.convert_to_i8();
}
pub fn finalize_weights(&mut self) {
self.w_q.finalize_weights();
self.w_k.finalize_weights();
self.w_v.finalize_weights();
self.w_o.finalize_weights();
}
pub fn finalize_weights_encoder(&mut self) {
self.w_q.finalize_weights_encoder();
self.w_k.finalize_weights_encoder();
self.w_v.finalize_weights_encoder();
self.w_o.finalize_weights_encoder();
}
#[must_use]
pub fn is_finalized(&self) -> bool {
self.w_q.is_finalized()
&& self.w_k.is_finalized()
&& self.w_v.is_finalized()
&& self.w_o.is_finalized()
}
pub fn w_q_mut(&mut self) -> &mut LinearWeights {
&mut self.w_q
}
pub fn w_k_mut(&mut self) -> &mut LinearWeights {
&mut self.w_k
}
pub fn w_v_mut(&mut self) -> &mut LinearWeights {
&mut self.w_v
}
pub fn w_o_mut(&mut self) -> &mut LinearWeights {
&mut self.w_o
}
pub fn fuse_qkv_weights(&mut self) {
if let (Some(wq), Some(wk), Some(wv)) = (
self.w_q.weight_f16().map(|s| s.to_vec()),
self.w_k.weight_f16().map(|s| s.to_vec()),
self.w_v.weight_f16().map(|s| s.to_vec()),
) {
let mut fused = Vec::with_capacity(wq.len() + wk.len() + wv.len());
fused.extend_from_slice(&wq);
fused.extend_from_slice(&wk);
fused.extend_from_slice(&wv);
self.w_qkv_f16 = Some(fused);
}
let mut b = Vec::with_capacity(self.d_model * 3);
b.extend_from_slice(&self.w_q.bias);
b.extend_from_slice(&self.w_k.bias);
b.extend_from_slice(&self.w_v.bias);
self.b_qkv = Some(b);
}
#[must_use]
pub fn has_fused_qkv(&self) -> bool {
self.w_qkv_f16.is_some() || self.w_q.is_i8()
}
pub fn forward_qkv_into(&self, input: &[f32], qkv_out: &mut [f32]) -> WhisperResult<()> {
let d = self.d_model;
debug_assert_eq!(qkv_out.len(), 3 * d);
if let (Some(w_qkv), Some(b_qkv)) = (&self.w_qkv_f16, &self.b_qkv) {
simd::tiled_matvec_f16_into(w_qkv, input, qkv_out, 3 * d, d);
simd::broadcast_add_inplace(qkv_out, b_qkv, 1, 3 * d);
Ok(())
} else {
self.w_q.forward_simd_into(input, 1, &mut qkv_out[..d])?;
self.w_k
.forward_simd_into(input, 1, &mut qkv_out[d..2 * d])?;
self.w_v
.forward_simd_into(input, 1, &mut qkv_out[2 * d..3 * d])?;
Ok(())
}
}
pub fn forward_qkv_batch_into(
&self,
input: &[f32],
batch_size: usize,
qkv_out: &mut [f32],
) -> WhisperResult<()> {
let d = self.d_model;
debug_assert_eq!(input.len(), batch_size * d);
debug_assert_eq!(qkv_out.len(), batch_size * 3 * d);
if let (Some(w_qkv), Some(b_qkv)) = (&self.w_qkv_f16, &self.b_qkv) {
const FP16_MATVEC_BATCH_LIMIT: usize = 8;
if batch_size <= FP16_MATVEC_BATCH_LIMIT {
for i in 0..batch_size {
let in_start = i * d;
let out_start = i * 3 * d;
simd::tiled_matvec_f16_into(
w_qkv,
&input[in_start..in_start + d],
&mut qkv_out[out_start..out_start + 3 * d],
3 * d,
d,
);
}
} else {
let mut buf = vec![0.0_f32; w_qkv.len()];
simd::dequant_f16_row(w_qkv, &mut buf);
let weight_t = simd::transpose(&buf, 3 * d, d);
let tmp = simd::matmul(input, &weight_t, batch_size, d, 3 * d);
qkv_out.copy_from_slice(&tmp);
}
simd::broadcast_add_inplace(qkv_out, b_qkv, batch_size, 3 * d);
Ok(())
} else {
let mut q = vec![0.0; batch_size * d];
let mut k = vec![0.0; batch_size * d];
let mut v = vec![0.0; batch_size * d];
self.w_q.forward_simd_into(input, batch_size, &mut q)?;
self.w_k.forward_simd_into(input, batch_size, &mut k)?;
self.w_v.forward_simd_into(input, batch_size, &mut v)?;
for i in 0..batch_size {
let out_start = i * 3 * d;
let in_start = i * d;
qkv_out[out_start..out_start + d].copy_from_slice(&q[in_start..in_start + d]);
qkv_out[out_start + d..out_start + 2 * d].copy_from_slice(&k[in_start..in_start + d]);
qkv_out[out_start + 2 * d..out_start + 3 * d].copy_from_slice(&v[in_start..in_start + d]);
}
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_linear_weights_new() {
let linear = LinearWeights::new(64, 128);
assert_eq!(linear.in_features, 64);
assert_eq!(linear.out_features, 128);
assert_eq!(linear.weight.len(), 128 * 64);
assert_eq!(linear.bias.len(), 128);
}
#[test]
fn test_linear_forward_identity() {
let mut linear = LinearWeights::new(4, 4);
for i in 0..4 {
linear.weight[i * 4 + i] = 1.0;
}
let input = vec![1.0, 2.0, 3.0, 4.0];
let output = linear.forward(&input, 1).expect("forward should succeed");
assert_eq!(output.len(), 4);
for i in 0..4 {
assert!(
(output[i] - input[i]).abs() < 1e-5,
"Identity should preserve input"
);
}
}
#[test]
fn test_linear_forward_with_bias() {
let mut linear = LinearWeights::new(2, 2);
linear.weight = vec![1.0, 0.0, 0.0, 1.0];
linear.bias = vec![1.0, 2.0];
let input = vec![3.0, 4.0];
let output = linear.forward(&input, 1).expect("forward should succeed");
assert!((output[0] - 4.0).abs() < 1e-5); assert!((output[1] - 6.0).abs() < 1e-5); }
#[test]
fn test_attention_new() {
let attn = MultiHeadAttention::new(8, 512);
assert_eq!(attn.n_heads(), 8);
assert_eq!(attn.d_model(), 512);
assert_eq!(attn.d_head(), 64);
}
#[test]
fn test_attention_scale() {
let attn = MultiHeadAttention::new(8, 512);
let expected_scale = 1.0 / (64.0_f32).sqrt();
assert!((attn.scale() - expected_scale).abs() < 1e-6);
}
#[test]
#[should_panic(expected = "must be divisible")]
fn test_attention_invalid_dimensions() {
let _ = MultiHeadAttention::new(8, 100); }
#[test]
fn test_scaled_dot_product_attention_basic() {
let attn = MultiHeadAttention::new(1, 4);
let query = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let value = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let output = attn
.scaled_dot_product_attention(&query, &key, &value, None)
.expect("attention should succeed");
assert_eq!(output.len(), 8); }
#[test]
fn test_scaled_dot_product_attention_with_mask() {
let attn = MultiHeadAttention::new(1, 4);
let query = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0];
let value = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let mask = vec![0.0, f32::NEG_INFINITY, 0.0, 0.0];
let output = attn
.scaled_dot_product_attention(&query, &key, &value, Some(&mask))
.expect("attention should succeed");
assert_eq!(output.len(), 8);
}
#[test]
fn test_attention_softmax_sums_to_one() {
let attn = MultiHeadAttention::new(1, 4);
let query = vec![1.0, 2.0, 3.0, 4.0]; let key = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0]; let value = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let _output = attn
.scaled_dot_product_attention(&query, &key, &value, None)
.expect("attention should succeed");
}
#[test]
fn test_causal_mask_shape() {
let mask = MultiHeadAttention::causal_mask(4);
assert_eq!(mask.len(), 16); }
#[test]
fn test_causal_mask_values() {
let mask = MultiHeadAttention::causal_mask(3);
assert_eq!(mask[0], 0.0); assert_eq!(mask[1], f32::NEG_INFINITY); assert_eq!(mask[2], f32::NEG_INFINITY);
assert_eq!(mask[3], 0.0); assert_eq!(mask[4], 0.0); assert_eq!(mask[5], f32::NEG_INFINITY);
assert_eq!(mask[6], 0.0); assert_eq!(mask[7], 0.0); assert_eq!(mask[8], 0.0); }
#[test]
fn test_forward_basic() {
let attn = MultiHeadAttention::new(2, 8);
let input = vec![0.0_f32; 16];
let output = attn.forward(&input, None).expect("forward should succeed");
assert_eq!(output.len(), 16); }
#[test]
fn test_forward_cross_basic() {
let attn = MultiHeadAttention::new(2, 8);
let x = vec![0.0_f32; 16]; let context = vec![0.0_f32; 24];
let output = attn
.forward_cross(&x, &context, None)
.expect("forward_cross should succeed");
assert_eq!(output.len(), 16); }
#[test]
fn test_forward_with_causal_mask() {
let attn = MultiHeadAttention::new(2, 8);
let input = vec![0.0_f32; 16]; let mask = MultiHeadAttention::causal_mask(2);
let output = attn
.forward(&input, Some(&mask))
.expect("forward should succeed");
assert_eq!(output.len(), 16);
}
#[test]
fn test_extract_head() {
let attn = MultiHeadAttention::new(2, 8);
let tensor: Vec<f32> = (0..16).map(|i| i as f32).collect();
let head0 = attn.extract_head(&tensor, 2, 0);
let head1 = attn.extract_head(&tensor, 2, 1);
assert_eq!(head0.len(), 8); assert_eq!(head1.len(), 8);
assert_eq!(head0[0..4], [0.0, 1.0, 2.0, 3.0]);
assert_eq!(head1[0..4], [4.0, 5.0, 6.0, 7.0]);
}
#[test]
fn test_concat_heads() {
let attn = MultiHeadAttention::new(2, 8);
let head0 = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]; let head1 = vec![9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0];
let concat = attn.concat_heads(&[head0, head1], 2);
assert_eq!(concat.len(), 16);
assert_eq!(concat[0..8], [1.0, 2.0, 3.0, 4.0, 9.0, 10.0, 11.0, 12.0]);
assert_eq!(concat[8..16], [5.0, 6.0, 7.0, 8.0, 13.0, 14.0, 15.0, 16.0]);
}
#[test]
fn test_extract_concat_roundtrip() {
let attn = MultiHeadAttention::new(2, 8);
let original: Vec<f32> = (0..16).map(|i| i as f32).collect();
let head0 = attn.extract_head(&original, 2, 0);
let head1 = attn.extract_head(&original, 2, 1);
let reconstructed = attn.concat_heads(&[head0, head1], 2);
assert_eq!(original, reconstructed);
}
#[test]
fn test_attention_size_mismatch() {
let attn = MultiHeadAttention::new(2, 8);
let query = vec![0.0_f32; 8]; let key = vec![0.0_f32; 8];
let value = vec![0.0_f32; 12];
let result = attn.scaled_dot_product_attention(&query, &key, &value, None);
assert!(result.is_err());
}
#[test]
fn test_forward_size_mismatch() {
let attn = MultiHeadAttention::new(2, 8);
let input = vec![0.0_f32; 15];
let result = attn.forward(&input, None);
assert!(result.is_err());
}
#[test]
fn test_weight_accessors() {
let attn = MultiHeadAttention::new(4, 64);
assert_eq!(attn.w_q().in_features, 64);
assert_eq!(attn.w_k().in_features, 64);
assert_eq!(attn.w_v().in_features, 64);
assert_eq!(attn.w_o().in_features, 64);
}
#[test]
fn test_linear_set_weight() {
let mut linear = LinearWeights::new(4, 4);
let weights = vec![1.0_f32; 16];
linear.set_weight(&weights);
assert!((linear.weight[0] - 1.0).abs() < f32::EPSILON);
assert!((linear.weight[15] - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_linear_set_bias() {
let mut linear = LinearWeights::new(4, 4);
let biases = vec![0.5_f32; 4];
linear.set_bias(&biases);
assert!((linear.bias[0] - 0.5).abs() < f32::EPSILON);
assert!((linear.bias[3] - 0.5).abs() < f32::EPSILON);
}
#[test]
fn test_linear_set_weight_partial() {
let mut linear = LinearWeights::new(4, 4);
let weights = vec![2.0_f32; 8];
linear.set_weight(&weights);
assert!((linear.weight[0] - 2.0).abs() < f32::EPSILON);
assert!((linear.weight[7] - 2.0).abs() < f32::EPSILON);
assert!((linear.weight[8] - 0.0).abs() < f32::EPSILON); }
#[test]
fn test_attention_set_query_weight() {
let mut attn = MultiHeadAttention::new(2, 8);
let weights = vec![1.0_f32; 64];
attn.set_query_weight(&weights);
assert!((attn.w_q().weight[0] - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_attention_set_key_weight() {
let mut attn = MultiHeadAttention::new(2, 8);
let weights = vec![2.0_f32; 64];
attn.set_key_weight(&weights);
assert!((attn.w_k().weight[0] - 2.0).abs() < f32::EPSILON);
}
#[test]
fn test_attention_set_value_weight() {
let mut attn = MultiHeadAttention::new(2, 8);
let weights = vec![3.0_f32; 64];
attn.set_value_weight(&weights);
assert!((attn.w_v().weight[0] - 3.0).abs() < f32::EPSILON);
}
#[test]
fn test_attention_set_out_weight() {
let mut attn = MultiHeadAttention::new(2, 8);
let weights = vec![4.0_f32; 64];
attn.set_out_weight(&weights);
assert!((attn.w_o().weight[0] - 4.0).abs() < f32::EPSILON);
}
#[test]
fn test_linear_forward_simd_identity() {
let mut linear = LinearWeights::new(4, 4);
for i in 0..4 {
linear.weight[i * 4 + i] = 1.0;
}
let input = vec![1.0, 2.0, 3.0, 4.0];
let output = linear
.forward_simd(&input, 1)
.expect("forward_simd should succeed");
assert_eq!(output.len(), 4);
for i in 0..4 {
assert!(
(output[i] - input[i]).abs() < 1e-4,
"SIMD Identity should preserve input: got {} expected {}",
output[i],
input[i]
);
}
}
#[test]
fn test_linear_forward_simd_with_bias() {
let mut linear = LinearWeights::new(2, 2);
linear.weight = vec![1.0, 0.0, 0.0, 1.0];
linear.bias = vec![1.0, 2.0];
let input = vec![3.0, 4.0];
let output = linear
.forward_simd(&input, 1)
.expect("forward_simd should succeed");
assert!(
(output[0] - 4.0).abs() < 1e-4,
"expected 4.0, got {}",
output[0]
); assert!(
(output[1] - 6.0).abs() < 1e-4,
"expected 6.0, got {}",
output[1]
); }
#[test]
fn test_linear_forward_simd_batch() {
let mut linear = LinearWeights::new(2, 2);
linear.weight = vec![2.0, 0.0, 0.0, 3.0];
linear.bias = vec![0.0, 0.0];
let input = vec![1.0, 2.0, 3.0, 4.0];
let output = linear
.forward_simd(&input, 2)
.expect("forward_simd should succeed");
assert_eq!(output.len(), 4);
assert!((output[0] - 2.0).abs() < 1e-4); assert!((output[1] - 6.0).abs() < 1e-4); assert!((output[2] - 6.0).abs() < 1e-4); assert!((output[3] - 12.0).abs() < 1e-4); }
#[test]
fn test_linear_forward_consistency() {
let mut linear = LinearWeights::new(4, 4);
for i in 0..16 {
linear.weight[i] = (i as f32) * 0.1;
}
linear.bias = vec![0.1, 0.2, 0.3, 0.4];
let input = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let output_regular = linear.forward(&input, 2).expect("forward should succeed");
let output_simd = linear
.forward_simd(&input, 2)
.expect("forward_simd should succeed");
assert_eq!(output_regular.len(), output_simd.len());
for i in 0..output_regular.len() {
assert!(
(output_regular[i] - output_simd[i]).abs() < 1e-3,
"SIMD and regular forward should match at index {}: {} vs {}",
i,
output_regular[i],
output_simd[i]
);
}
}
#[test]
fn test_linear_finalize_weights() {
let mut linear = LinearWeights::new(4, 4);
for i in 0..16 {
linear.weight[i] = (i as f32) * 0.1;
}
linear.bias = vec![0.1, 0.2, 0.3, 0.4];
assert!(!linear.is_finalized());
linear.finalize_weights();
assert!(linear.is_finalized());
let input = vec![1.0, 2.0, 3.0, 4.0];
let output_finalized = linear
.forward_simd(&input, 1)
.expect("forward_simd should succeed");
linear.invalidate_cache();
assert!(!linear.is_finalized());
let output_unfinalized = linear
.forward_simd(&input, 1)
.expect("forward_simd should succeed");
assert_eq!(output_finalized.len(), output_unfinalized.len());
for i in 0..output_finalized.len() {
assert!(
(output_finalized[i] - output_unfinalized[i]).abs() < 1e-6,
"Finalized and unfinalized should match"
);
}
}
#[test]
fn test_attention_finalize_weights() {
let mut attn = MultiHeadAttention::new(2, 8);
assert!(!attn.is_finalized());
attn.finalize_weights();
assert!(attn.is_finalized());
assert!(attn.w_q().is_finalized());
assert!(attn.w_k().is_finalized());
assert!(attn.w_v().is_finalized());
assert!(attn.w_o().is_finalized());
}
#[test]
fn test_scaled_dot_product_attention_simd_basic() {
let attn = MultiHeadAttention::new(1, 4);
let query = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let value = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let output = attn
.scaled_dot_product_attention_simd(&query, &key, &value, None)
.expect("SIMD attention should succeed");
assert_eq!(output.len(), 8); }
#[test]
fn test_scaled_dot_product_attention_simd_with_mask() {
let attn = MultiHeadAttention::new(1, 4);
let query = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0];
let value = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let mask = vec![0.0, f32::NEG_INFINITY, 0.0, 0.0];
let output = attn
.scaled_dot_product_attention_simd(&query, &key, &value, Some(&mask))
.expect("SIMD attention with mask should succeed");
assert_eq!(output.len(), 8);
}
#[test]
fn test_attention_simd_error_handling() {
let attn = MultiHeadAttention::new(2, 8);
let query = vec![0.0_f32; 9]; let key = vec![0.0_f32; 8];
let value = vec![0.0_f32; 8];
let result = attn.scaled_dot_product_attention_simd(&query, &key, &value, None);
assert!(result.is_err());
}
#[test]
fn test_flash_attention_basic() {
let query = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let value = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let config = FlashAttentionConfig::new(2, 2, 4, 2);
let output = flash_attention(&query, &key, &value, config, None);
assert_eq!(output.len(), 8);
for &v in &output {
assert!(v.is_finite(), "Flash attention output should be finite");
}
}
#[test]
fn test_flash_attention_matches_standard() {
let attn = MultiHeadAttention::new(1, 4);
let query = vec![1.0, 0.5, 0.0, 0.0, 0.5, 1.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.5, 0.5, 0.0, 0.0];
let value = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 2.0, 3.0, 4.0, 5.0];
let standard = attn
.scaled_dot_product_attention(&query, &key, &value, None)
.expect("standard attention");
let config = FlashAttentionConfig::new(2, 3, 4, 2);
let flash = flash_attention(&query, &key, &value, config, None);
assert_eq!(standard.len(), flash.len());
for i in 0..standard.len() {
assert!(
(standard[i] - flash[i]).abs() < 1e-4,
"Flash attention should match standard at index {}: {} vs {}",
i,
standard[i],
flash[i]
);
}
}
#[test]
fn test_flash_attention_with_mask() {
let query = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0];
let value = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let mask = vec![0.0, f32::NEG_INFINITY, 0.0, 0.0];
let config = FlashAttentionConfig::new(2, 2, 4, 2);
let output = flash_attention(&query, &key, &value, config, Some(&mask));
assert_eq!(output.len(), 8);
assert!(
(output[0] - 1.0).abs() < 1e-4,
"First output[0] should be 1.0"
);
}
#[test]
fn test_flash_attention_simd_matches_scalar() {
let query = vec![1.0, 0.5, 0.0, 0.0, 0.5, 1.0, 0.0, 0.0];
let key = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let value = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let config = FlashAttentionConfig::new(2, 2, 4, 2);
let scalar = flash_attention(&query, &key, &value, config, None);
let simd = flash_attention_simd(&query, &key, &value, config, None);
assert_eq!(scalar.len(), simd.len());
for i in 0..scalar.len() {
assert!(
(scalar[i] - simd[i]).abs() < 1e-5,
"SIMD flash attention should match scalar at index {}: {} vs {}",
i,
scalar[i],
simd[i]
);
}
}
#[test]
fn test_flash_attention_different_block_sizes() {
let query = vec![1.0; 32]; let key = vec![1.0; 32];
let value = vec![1.0; 32];
let config_2 = FlashAttentionConfig::new(8, 8, 4, 2);
let config_4 = FlashAttentionConfig::new(8, 8, 4, 4);
let config_8 = FlashAttentionConfig::new(8, 8, 4, 8);
let out_block_2 = flash_attention(&query, &key, &value, config_2, None);
let out_block_4 = flash_attention(&query, &key, &value, config_4, None);
let out_block_8 = flash_attention(&query, &key, &value, config_8, None);
for i in 0..out_block_2.len() {
assert!(
(out_block_2[i] - out_block_4[i]).abs() < 1e-5,
"Block size 2 vs 4 mismatch at {i}"
);
assert!(
(out_block_4[i] - out_block_8[i]).abs() < 1e-5,
"Block size 4 vs 8 mismatch at {i}"
);
}
}
#[test]
fn test_forward_cross_flash() {
let attn = MultiHeadAttention::new(2, 8);
let x = vec![0.1_f32; 32];
let context = vec![0.2_f32; 32];
let output = attn
.forward_cross_flash(&x, &context, None, FLASH_ATTENTION_BLOCK_SIZE)
.expect("forward_cross_flash");
assert_eq!(output.len(), 32); }
#[test]
#[cfg(feature = "realizar-inference")]
fn test_forward_cross_flash_v2() {
let attn = MultiHeadAttention::new(2, 8);
let x = vec![0.1_f32; 32];
let context = vec![0.2_f32; 32];
let output = attn
.forward_cross_flash_v2(&x, &context, None)
.expect("forward_cross_flash_v2");
assert_eq!(output.len(), 32);
let output_v1 = attn
.forward_cross_flash(&x, &context, None, FLASH_ATTENTION_BLOCK_SIZE)
.expect("forward_cross_flash");
for i in 0..output.len() {
assert!(
(output[i] - output_v1[i]).abs() < 1e-3,
"FlashAttention-2 should match custom at index {}: {} vs {}",
i,
output[i],
output_v1[i]
);
}
}
#[test]
#[cfg(feature = "realizar-inference")]
fn test_forward_cross_optimal() {
let attn = MultiHeadAttention::new(2, 8);
let x_short = vec![0.1_f32; 32]; let ctx_short = vec![0.2_f32; 32];
let output = attn
.forward_cross_optimal(&x_short, &ctx_short, None)
.expect("forward_cross_optimal");
assert_eq!(output.len(), 32);
let x_long = vec![0.1_f32; 1024 * 8]; let ctx_long = vec![0.2_f32; 1024 * 8];
let output_long = attn
.forward_cross_optimal(&x_long, &ctx_long, None)
.expect("forward_cross_optimal long");
assert_eq!(output_long.len(), 1024 * 8);
}
#[test]
fn test_forward_cross_auto() {
let attn = MultiHeadAttention::new(2, 8);
let x_short = vec![0.1_f32; 32]; let ctx_short = vec![0.2_f32; 32];
let output_short = attn
.forward_cross_auto(&x_short, &ctx_short, None)
.expect("forward_cross_auto short");
assert_eq!(output_short.len(), 32);
let output_standard = attn
.forward_cross(&x_short, &ctx_short, None)
.expect("forward_cross standard");
for i in 0..output_short.len() {
assert!(
(output_short[i] - output_standard[i]).abs() < 1e-3,
"Auto should match standard for short sequences at {i}"
);
}
}
#[test]
fn test_forward_streaming_no_cache() {
let attn = MultiHeadAttention::new(2, 8);
let x = vec![0.1_f32; 8]; let empty_key: Vec<f32> = vec![];
let empty_value: Vec<f32> = vec![];
let (output, new_k, new_v) = attn
.forward_streaming(&x, &empty_key, &empty_value, None)
.expect("forward_streaming no cache");
assert_eq!(output.len(), 8);
assert_eq!(new_k.len(), 8); assert_eq!(new_v.len(), 8); }
#[test]
fn test_forward_streaming_with_cache() {
let attn = MultiHeadAttention::new(2, 8);
let x1 = vec![0.1_f32; 8];
let (_, k1, v1) = attn
.forward_streaming(&x1, &[], &[], None)
.expect("first streaming call");
let x2 = vec![0.2_f32; 8];
let (output2, k2, v2) = attn
.forward_streaming(&x2, &k1, &v1, None)
.expect("second streaming call");
assert_eq!(output2.len(), 8);
assert_eq!(k2.len(), 8); assert_eq!(v2.len(), 8); }
#[test]
fn test_forward_streaming_multi_token() {
let attn = MultiHeadAttention::new(2, 8);
let x = vec![0.1_f32; 24]; let (output, new_k, new_v) = attn
.forward_streaming(&x, &[], &[], None)
.expect("forward_streaming multi-token");
assert_eq!(output.len(), 24);
assert_eq!(new_k.len(), 24);
assert_eq!(new_v.len(), 24);
}
#[test]
fn test_forward_self_streaming() {
let attn = MultiHeadAttention::new(2, 8);
let x = vec![0.1_f32; 8];
let (output, new_k, new_v) = attn
.forward_self_streaming(&x, &[], &[])
.expect("forward_self_streaming");
assert_eq!(output.len(), 8);
assert_eq!(new_k.len(), 8);
assert_eq!(new_v.len(), 8);
}
#[test]
fn regression_flash_attention_accuracy() {
let attn = MultiHeadAttention::new(1, 64);
let seq_len = 32;
let d_head = 64;
let q: Vec<f32> = (0..seq_len * d_head)
.map(|i| (i as f32 * 0.1).sin() * 0.5)
.collect();
let k: Vec<f32> = (0..seq_len * d_head)
.map(|i| (i as f32 * 0.2).cos() * 0.5)
.collect();
let v: Vec<f32> = (0..seq_len * d_head)
.map(|i| (i as f32 * 0.3).sin() * 0.5)
.collect();
let standard = attn
.scaled_dot_product_attention(&q, &k, &v, None)
.expect("standard attention");
let config = FlashAttentionConfig::with_default_block_size(seq_len, seq_len, d_head);
let flash_scalar = flash_attention(&q, &k, &v, config, None);
let flash_simd = flash_attention_simd(&q, &k, &v, config, None);
let tolerance = 1e-4;
for i in 0..standard.len() {
assert!(
(standard[i] - flash_scalar[i]).abs() < tolerance,
"Flash scalar mismatch at {i}: {} vs {}",
standard[i],
flash_scalar[i]
);
assert!(
(standard[i] - flash_simd[i]).abs() < tolerance,
"Flash SIMD mismatch at {i}: {} vs {}",
standard[i],
flash_simd[i]
);
}
}
#[test]
fn regression_streaming_attention_consistency() {
let attn = MultiHeadAttention::new(2, 16);
let d_model = 16;
let mut cached_k = Vec::new();
let mut cached_v = Vec::new();
let mut all_outputs = Vec::new();
for i in 0..5 {
let x: Vec<f32> = (0..d_model)
.map(|j| (i * d_model + j) as f32 * 0.01)
.collect();
let (output, new_k, new_v) = attn
.forward_streaming(&x, &cached_k, &cached_v, None)
.expect("streaming attention");
cached_k.extend_from_slice(&new_k);
cached_v.extend_from_slice(&new_v);
all_outputs.push(output);
}
for (i, output) in all_outputs.iter().enumerate() {
assert_eq!(output.len(), d_model, "Token {i} output size mismatch");
}
assert_eq!(cached_k.len(), 5 * d_model);
assert_eq!(cached_v.len(), 5 * d_model);
}
#[test]
fn regression_simd_numerical_stability() {
use crate::simd;
let large_vals: Vec<f32> = (0..256).map(|i| 100.0 + i as f32 * 0.1).collect();
let small_vals: Vec<f32> = (0..256).map(|i| 1e-6 + i as f32 * 1e-8).collect();
let mixed_vals: Vec<f32> = (0..256)
.map(|i| if i % 2 == 0 { 100.0 } else { 0.001 })
.collect();
let softmax_large = simd::softmax(&large_vals);
assert!(
softmax_large.iter().all(|&x| x.is_finite()),
"Softmax produced non-finite values for large inputs"
);
let sum: f32 = softmax_large.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "Softmax sum not 1.0: {}", sum);
let softmax_small = simd::softmax(&small_vals);
assert!(
softmax_small.iter().any(|&x| x > 0.0),
"Softmax underflowed to all zeros"
);
let dot_mixed = simd::dot(&mixed_vals, &mixed_vals);
assert!(
dot_mixed.is_finite(),
"Dot product produced non-finite value"
);
}
#[test]
fn regression_flash_attention_block_size_invariance() {
let seq_len = 64;
let d_head = 32;
let q: Vec<f32> = (0..seq_len * d_head)
.map(|i| (i as f32 * 0.1).sin())
.collect();
let k: Vec<f32> = (0..seq_len * d_head)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..seq_len * d_head)
.map(|i| (i as f32 * 0.15).sin())
.collect();
let block_sizes = [8, 16, 32, 64];
let mut results = Vec::new();
for &block_size in &block_sizes {
let config = FlashAttentionConfig::new(seq_len, seq_len, d_head, block_size);
let output = flash_attention_simd(&q, &k, &v, config, None);
results.push(output);
}
let tolerance = 1e-5;
for i in 1..results.len() {
for j in 0..results[0].len() {
assert!(
(results[0][j] - results[i][j]).abs() < tolerance,
"Block size {} differs from block size {} at position {}: {} vs {}",
block_sizes[0],
block_sizes[i],
j,
results[0][j],
results[i][j]
);
}
}
}
#[test]
fn regression_multihead_output_shapes() {
let configs = [
(6, 384), (8, 512), (12, 768), ];
for (n_heads, d_model) in configs {
let attn = MultiHeadAttention::new(n_heads, d_model);
for seq_len in [1, 10, 100] {
let x = vec![0.1_f32; seq_len * d_model];
let context = vec![0.2_f32; seq_len * d_model];
let output = attn
.forward_cross(&x, &context, None)
.expect("forward_cross");
assert_eq!(
output.len(),
seq_len * d_model,
"Output shape mismatch for n_heads={}, d_model={}, seq_len={}",
n_heads,
d_model,
seq_len
);
}
}
}
#[test]
fn tdd_simd_dispatch_forward_cross_matches_scalar() {
let mut attn = MultiHeadAttention::new(6, 384); attn.finalize_weights();
let seq_len = 10;
let d_model = 384;
let tolerance = 1e-4;
let x: Vec<f32> = (0..seq_len * d_model)
.map(|i| (i as f32 * 0.01).sin())
.collect();
let context: Vec<f32> = (0..seq_len * d_model)
.map(|i| (i as f32 * 0.02).cos())
.collect();
let scalar_output = attn
.forward_cross(&x, &context, None)
.expect("scalar forward_cross");
let simd_output = attn
.forward_cross_simd(&x, &context, None)
.expect("simd forward_cross");
assert_eq!(
scalar_output.len(),
simd_output.len(),
"Output shapes must match"
);
let mut max_diff: f32 = 0.0;
for (i, (s, simd)) in scalar_output.iter().zip(simd_output.iter()).enumerate() {
let diff = (s - simd).abs();
max_diff = max_diff.max(diff);
assert!(
diff < tolerance,
"SIMD mismatch at index {}: scalar={}, simd={}, diff={}",
i,
s,
simd,
diff
);
}
println!(
"SIMD accuracy test passed: max_diff={:.2e} (tolerance={:.2e})",
max_diff, tolerance
);
}
#[test]
fn tdd_forward_uses_simd_dispatch() {
let mut attn = MultiHeadAttention::new(4, 64);
attn.finalize_weights();
let seq_len = 5;
let d_model = 64;
let tolerance = 1e-4;
let x: Vec<f32> = (0..seq_len * d_model)
.map(|i| (i as f32 * 0.1).sin())
.collect();
let forward_output = attn.forward(&x, None).expect("forward");
let simd_output = attn
.forward_cross_simd(&x, &x, None)
.expect("forward_cross_simd");
if cfg!(feature = "simd") {
for (i, (f, s)) in forward_output.iter().zip(simd_output.iter()).enumerate() {
let diff = (f - s).abs();
assert!(
diff < tolerance,
"forward() should match forward_cross_simd() when simd enabled. Index {}: {} vs {}, diff={}",
i, f, s, diff
);
}
println!("SIMD dispatch verified: forward() uses forward_cross_simd()");
}
}
#[test]
fn tdd_forward_cross_auto_dispatch() {
let mut attn = MultiHeadAttention::new(2, 32);
attn.finalize_weights();
let seq_len = 4;
let d_model = 32;
let x: Vec<f32> = (0..seq_len * d_model).map(|i| i as f32 * 0.01).collect();
let ctx: Vec<f32> = (0..seq_len * d_model).map(|i| i as f32 * 0.02).collect();
let output = attn
.forward_cross_dispatch(&x, &ctx, None)
.expect("forward_cross_dispatch");
assert_eq!(output.len(), seq_len * d_model);
println!(
"Auto-dispatch forward_cross completed with {} elements",
output.len()
);
}
fn make_fused_mha(d_model: usize) -> MultiHeadAttention {
let n_heads = if d_model >= 8 { d_model / 64.max(1) } else { 1 };
let adjusted_d = n_heads * (d_model / n_heads); let n_heads = if adjusted_d == 0 {
1
} else {
adjusted_d / (adjusted_d / n_heads)
};
let mut attn = MultiHeadAttention::new(n_heads, d_model);
let n = d_model * d_model;
let make_f16 = |seed: u32| -> Vec<u16> {
(0..n)
.map(|i| {
let v = ((i as f32 + seed as f32) * 0.001).sin() * 0.1;
half::f16::from_f32(v).to_bits()
})
.collect()
};
attn.set_query_weight_f16(&make_f16(1));
attn.set_key_weight_f16(&make_f16(7));
attn.set_value_weight_f16(&make_f16(13));
let make_bias = |seed: u32| -> Vec<f32> {
(0..d_model)
.map(|i| ((i as f32 + seed as f32) * 0.01).cos() * 0.05)
.collect()
};
attn.set_query_bias(&make_bias(2));
attn.set_key_bias(&make_bias(8));
attn.set_value_bias(&make_bias(14));
attn.fuse_qkv_weights();
attn
}
#[test]
fn pv_fused_qkv_equivalence() {
let d = 384; let attn = make_fused_mha(d);
assert!(attn.has_fused_qkv());
let input: Vec<f32> = (0..d).map(|i| (i as f32 * 0.01).sin()).collect();
let mut q_sep = vec![0.0f32; d];
let mut k_sep = vec![0.0f32; d];
let mut v_sep = vec![0.0f32; d];
attn.w_q().forward_simd_into(&input, 1, &mut q_sep).unwrap();
attn.w_k().forward_simd_into(&input, 1, &mut k_sep).unwrap();
attn.w_v().forward_simd_into(&input, 1, &mut v_sep).unwrap();
let mut qkv = vec![0.0f32; 3 * d];
attn.forward_qkv_into(&input, &mut qkv).unwrap();
for i in 0..d {
let diff_q = (q_sep[i] - qkv[i]).abs();
let diff_k = (k_sep[i] - qkv[d + i]).abs();
let diff_v = (v_sep[i] - qkv[2 * d + i]).abs();
assert!(
diff_q < 1e-4,
"q[{i}]: sep={}, fused={}, diff={diff_q}",
q_sep[i],
qkv[i]
);
assert!(
diff_k < 1e-4,
"k[{i}]: sep={}, fused={}, diff={diff_k}",
k_sep[i],
qkv[d + i]
);
assert!(
diff_v < 1e-4,
"v[{i}]: sep={}, fused={}, diff={diff_v}",
v_sep[i],
qkv[2 * d + i]
);
}
}
#[test]
fn pv_fused_qkv_output_dimension() {
for d in [64, 384, 512, 768] {
let n_heads = d / 64;
let mut attn = MultiHeadAttention::new(n_heads, d);
let n = d * d;
let zeros_f16: Vec<u16> = vec![0; n];
attn.set_query_weight_f16(&zeros_f16);
attn.set_key_weight_f16(&zeros_f16);
attn.set_value_weight_f16(&zeros_f16);
attn.fuse_qkv_weights();
let input = vec![1.0f32; d];
let mut qkv = vec![0.0f32; 3 * d];
attn.forward_qkv_into(&input, &mut qkv).unwrap();
assert_eq!(qkv.len(), 3 * d, "d_model={d}");
}
}
#[test]
fn pv_fused_qkv_weight_layout() {
let d = 64;
let attn = make_fused_mha(d);
let w_qkv = attn.w_qkv_f16.as_ref().unwrap();
let wq = attn.w_q.weight_f16().unwrap();
let wk = attn.w_k.weight_f16().unwrap();
let wv = attn.w_v.weight_f16().unwrap();
let n = d * d;
assert_eq!(w_qkv.len(), 3 * n);
assert_eq!(&w_qkv[..n], wq);
assert_eq!(&w_qkv[n..2 * n], wk);
assert_eq!(&w_qkv[2 * n..3 * n], wv);
}
#[test]
fn pv_fused_qkv_bias_layout() {
let d = 64;
let attn = make_fused_mha(d);
let b_qkv = attn.b_qkv.as_ref().unwrap();
assert_eq!(b_qkv.len(), 3 * d);
assert_eq!(&b_qkv[..d], attn.w_q.bias.as_slice());
assert_eq!(&b_qkv[d..2 * d], attn.w_k.bias.as_slice());
assert_eq!(&b_qkv[2 * d..3 * d], attn.w_v.bias.as_slice());
}
#[test]
fn pv_fused_qkv_whisper_dimensions() {
for d in [384, 512, 768, 1024, 1280] {
let n_heads = d / 64;
let attn = make_fused_mha(d);
assert!(attn.has_fused_qkv(), "d_model={d}");
let input: Vec<f32> = (0..d).map(|i| (i as f32 * 0.005).sin()).collect();
let mut qkv = vec![0.0f32; 3 * d];
attn.forward_qkv_into(&input, &mut qkv).unwrap();
let q_sum: f32 = qkv[..d].iter().map(|x| x.abs()).sum();
let k_sum: f32 = qkv[d..2 * d].iter().map(|x| x.abs()).sum();
let v_sum: f32 = qkv[2 * d..].iter().map(|x| x.abs()).sum();
assert!(q_sum > 0.0, "d={d}: Q all zero");
assert!(k_sum > 0.0, "d={d}: K all zero");
assert!(v_sum > 0.0, "d={d}: V all zero");
let _ = n_heads;
}
}
}