use anyhow::Result;
use crate::tensor_core::{Tensor, Device, DataType};
pub struct RoPECache {
cos: Tensor,
sin: Tensor,
head_dim: usize,
max_seq_len: usize,
base: f32,
}
impl RoPECache {
pub fn new(head_dim: usize, max_seq_len: usize, base: f32, device: &Device) -> Result<Self> {
let half_dim = head_dim / 2;
let mut inv_freq = vec![0.0f32; half_dim];
for i in 0..half_dim {
let exponent = (2.0 * i as f32) / head_dim as f32;
inv_freq[i] = 1.0 / base.powf(exponent);
}
let positions: Vec<f32> = (0..max_seq_len).map(|i| i as f32).collect();
let mut cos_data = vec![0.0f32; max_seq_len * half_dim];
let mut sin_data = vec![0.0f32; max_seq_len * half_dim];
for pos in 0..max_seq_len {
for freq_idx in 0..half_dim {
let angle = positions[pos] * inv_freq[freq_idx];
cos_data[pos * half_dim + freq_idx] = angle.cos();
sin_data[pos * half_dim + freq_idx] = angle.sin();
}
}
let cos = Tensor::from_f32_slice(&cos_data, &[max_seq_len, half_dim], device)?;
let sin = Tensor::from_f32_slice(&sin_data, &[max_seq_len, half_dim], device)?;
Ok(Self {
cos,
sin,
head_dim,
max_seq_len,
base,
})
}
pub fn get(&self, seq_len: usize) -> Result<(Tensor, Tensor)> {
if seq_len > self.max_seq_len {
return Err(anyhow::anyhow!(
"Requested seq_len {} exceeds max_seq_len {}",
seq_len,
self.max_seq_len
));
}
let cos = self.cos.narrow(0, 0, seq_len)?;
let sin = self.sin.narrow(0, 0, seq_len)?;
Ok((cos, sin))
}
pub fn get_range(&self, start: usize, end: usize) -> Result<(Tensor, Tensor)> {
if end > self.max_seq_len {
return Err(anyhow::anyhow!(
"Requested end {} exceeds max_seq_len {}",
end,
self.max_seq_len
));
}
if start >= end {
return Err(anyhow::anyhow!(
"Invalid range: start {} >= end {}",
start,
end
));
}
let len = end - start;
let cos = self.cos.narrow(0, start, len)?;
let sin = self.sin.narrow(0, start, len)?;
Ok((cos, sin))
}
pub fn head_dim(&self) -> usize {
self.head_dim
}
pub fn max_seq_len(&self) -> usize {
self.max_seq_len
}
pub fn base(&self) -> f32 {
self.base
}
}
pub struct CausalMaskCache {
mask: Tensor,
max_seq_len: usize,
}
impl CausalMaskCache {
pub fn new(max_seq_len: usize, device: &Device) -> Result<Self> {
let mut mask_data = vec![0.0f32; max_seq_len * max_seq_len];
for i in 0..max_seq_len {
for j in 0..max_seq_len {
if j > i {
mask_data[i * max_seq_len + j] = f32::NEG_INFINITY;
}
}
}
let mask = Tensor::from_f32_slice(&mask_data, &[max_seq_len, max_seq_len], device)?;
Ok(Self { mask, max_seq_len })
}
pub fn get(&self, seq_len: usize) -> Result<Tensor> {
if seq_len > self.max_seq_len {
return Err(anyhow::anyhow!(
"Requested seq_len {} exceeds max_seq_len {}",
seq_len,
self.max_seq_len
));
}
let mask = self.mask.narrow(0, 0, seq_len)?.narrow(1, 0, seq_len)?;
Ok(mask)
}
pub fn get_broadcast(&self, seq_len: usize) -> Result<Tensor> {
let mask = self.get(seq_len)?;
mask.reshape(&[1, 1, seq_len, seq_len])
}
pub fn max_seq_len(&self) -> usize {
self.max_seq_len
}
}
pub struct SlidingWindowMaskCache {
mask: Tensor,
max_seq_len: usize,
window_size: usize,
}
impl SlidingWindowMaskCache {
pub fn new(max_seq_len: usize, window_size: usize, device: &Device) -> Result<Self> {
let mut mask_data = vec![f32::NEG_INFINITY; max_seq_len * max_seq_len];
for i in 0..max_seq_len {
let start = if i >= window_size { i - window_size + 1 } else { 0 };
for j in start..=i {
mask_data[i * max_seq_len + j] = 0.0;
}
}
let mask = Tensor::from_f32_slice(&mask_data, &[max_seq_len, max_seq_len], device)?;
Ok(Self {
mask,
max_seq_len,
window_size,
})
}
pub fn get(&self, seq_len: usize) -> Result<Tensor> {
if seq_len > self.max_seq_len {
return Err(anyhow::anyhow!(
"Requested seq_len {} exceeds max_seq_len {}",
seq_len,
self.max_seq_len
));
}
let mask = self.mask.narrow(0, 0, seq_len)?.narrow(1, 0, seq_len)?;
Ok(mask)
}
pub fn get_broadcast(&self, seq_len: usize) -> Result<Tensor> {
let mask = self.get(seq_len)?;
mask.reshape(&[1, 1, seq_len, seq_len])
}
pub fn window_size(&self) -> usize {
self.window_size
}
pub fn max_seq_len(&self) -> usize {
self.max_seq_len
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rope_cache_creation() {
let rope = RoPECache::new(64, 128, 10000.0, &Device::CPU).unwrap();
assert_eq!(rope.head_dim(), 64);
assert_eq!(rope.max_seq_len(), 128);
}
#[test]
fn test_rope_cache_get() {
let rope = RoPECache::new(64, 128, 10000.0, &Device::CPU).unwrap();
let (cos, sin) = rope.get(32).unwrap();
assert_eq!(cos.shape(), &[32, 32]); assert_eq!(sin.shape(), &[32, 32]);
}
#[test]
fn test_causal_mask_cache() {
let cache = CausalMaskCache::new(64, &Device::CPU).unwrap();
let mask = cache.get(8).unwrap();
assert_eq!(mask.shape(), &[8, 8]);
}
#[test]
fn test_causal_mask_values() {
let cache = CausalMaskCache::new(4, &Device::CPU).unwrap();
let mask = cache.get(4).unwrap();
let mask_flat = mask.reshape(&[16]).unwrap();
let values = mask_flat.to_vec_f32().unwrap();
assert_eq!(values[0], 0.0);
assert!(values[1].is_infinite() && values[1] < 0.0);
assert_eq!(values[12], 0.0); assert_eq!(values[13], 0.0); assert_eq!(values[14], 0.0); assert_eq!(values[15], 0.0); }
#[test]
fn test_sliding_window_mask() {
let cache = SlidingWindowMaskCache::new(8, 3, &Device::CPU).unwrap();
let mask = cache.get(8).unwrap();
let mask_flat = mask.reshape(&[64]).unwrap();
let values = mask_flat.to_vec_f32().unwrap();
assert!(values[5 * 8 + 2].is_infinite()); assert_eq!(values[5 * 8 + 3], 0.0); assert_eq!(values[5 * 8 + 4], 0.0); assert_eq!(values[5 * 8 + 5], 0.0); }
}