use candle_core::{D, DType, Device, IndexOp, Tensor};
use crate::error::{CandleOcrError, Result};
fn split_tensor(t: &Tensor, splits: &[usize], dim: usize) -> Result<Vec<Tensor>> {
let mut results: Vec<Tensor> = Vec::with_capacity(splits.len());
let mut offset = 0;
for &size in splits {
results.push(t.narrow(dim, offset, size)?);
offset += size;
}
Ok(results)
}
pub fn compute_default_rope_parameters(dim: usize, base: f32) -> Vec<f32> {
(0..dim)
.step_by(2)
.map(|i| 1.0_f32 / base.powf(i as f32 / dim as f32))
.collect()
}
pub fn rotate_half(x: &Tensor) -> Result<Tensor> {
let half_dim = x.dim(D::Minus1)? / 2;
let x1 = x.narrow(D::Minus1, 0, half_dim)?;
let x2 = x.narrow(D::Minus1, half_dim, half_dim)?;
let x2_neg = x2.affine(-1.0, 0.0)?;
Ok(Tensor::cat(&[&x2_neg, &x1], D::Minus1)?.contiguous()?)
}
pub fn apply_rotary_pos_emb(
q: &Tensor,
k: &Tensor,
cos: &Tensor,
sin: &Tensor,
tof32: bool,
) -> Result<(Tensor, Tensor)> {
let mut cos = cos.clone();
let mut sin = sin.clone();
if cos.rank() == 2 {
cos = cos.unsqueeze(0)?.unsqueeze(0)?;
sin = sin.unsqueeze(0)?.unsqueeze(0)?;
}
if cos.rank() == 3 {
cos = cos.unsqueeze(1)?;
sin = sin.unsqueeze(1)?;
}
let orig_dtype = q.dtype();
let q_f = if tof32 { q.to_dtype(DType::F32)? } else { q.clone() };
let k_f = if tof32 { k.to_dtype(DType::F32)? } else { k.clone() };
let cos = cos.to_dtype(q_f.dtype())?;
let sin = sin.to_dtype(q_f.dtype())?;
let q_embed = q_f
.broadcast_mul(&cos)?
.add(&rotate_half(&q_f)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
let k_embed = k_f
.broadcast_mul(&cos)?
.add(&rotate_half(&k_f)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
Ok((q_embed, k_embed))
}
pub fn apply_rotary_pos_emb_vision(q: &Tensor, k: &Tensor, cos: &Tensor, sin: &Tensor) -> Result<(Tensor, Tensor)> {
let cos = cos.unsqueeze(D::Minus2)?.to_dtype(q.dtype())?;
let sin = sin.unsqueeze(D::Minus2)?.to_dtype(q.dtype())?;
let q_embed = q.broadcast_mul(&cos)?.add(&rotate_half(q)?.broadcast_mul(&sin)?)?;
let k_embed = k.broadcast_mul(&cos)?.add(&rotate_half(k)?.broadcast_mul(&sin)?)?;
Ok((q_embed, k_embed))
}
pub fn apply_rotary_pos_emb_roformer(q: &Tensor, k: &Tensor, cos: &Tensor, sin: &Tensor) -> Result<(Tensor, Tensor)> {
let ori_dtype = q.dtype();
let (bs, n_head, seq_len, dim) = q.dims4()?;
let half_dim = dim / 2;
let rotr = cos.narrow(D::Minus1, 0, half_dim)?.to_dtype(DType::F32)?;
let roti = sin.narrow(D::Minus1, 0, half_dim)?.to_dtype(DType::F32)?;
let q_f = q.reshape((bs, n_head, seq_len, half_dim, 2))?.to_dtype(DType::F32)?;
let qr = q_f.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let qi = q_f.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?;
let k_f = k.reshape((bs, n_head, seq_len, half_dim, 2))?.to_dtype(DType::F32)?;
let kr = k_f.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let ki = k_f.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?;
let qor = qr.broadcast_mul(&rotr)?.sub(&qi.broadcast_mul(&roti)?)?;
let qoi = qr.broadcast_mul(&roti)?.add(&qi.broadcast_mul(&rotr)?)?;
let kor = kr.broadcast_mul(&rotr)?.sub(&ki.broadcast_mul(&roti)?)?;
let koi = kr.broadcast_mul(&roti)?.add(&ki.broadcast_mul(&rotr)?)?;
let q_out = Tensor::stack(&[qor, qoi], D::Minus1)?
.reshape((bs, n_head, seq_len, dim))?
.to_dtype(ori_dtype)?;
let k_out = Tensor::stack(&[kor, koi], D::Minus1)?
.reshape((bs, n_head, seq_len, dim))?
.to_dtype(ori_dtype)?;
Ok((q_out, k_out))
}
#[derive(Debug, Clone)]
pub struct RoPE {
inv_freq: Tensor,
}
impl RoPE {
pub fn new(dim: usize, theta_base: f32, device: &Device) -> Result<Self> {
let inv_freq = compute_default_rope_parameters(dim, theta_base);
let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?;
Ok(Self { inv_freq })
}
pub fn forward(&self, seqlen_offset: usize, seq_len: usize, device: &Device) -> Result<(Tensor, Tensor)> {
let positions = Tensor::arange(
seqlen_offset as f32,
(seqlen_offset + seq_len) as f32,
self.inv_freq.device(),
)?
.reshape((seq_len, 1))?;
let freqs = positions.matmul(&self.inv_freq)?;
let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?
.contiguous()?
.to_device(device)?;
Ok((emb.cos()?, emb.sin()?))
}
pub fn forward_repeat_interleave(
&self,
seqlen_offset: usize,
seq_len: usize,
device: &Device,
) -> Result<(Tensor, Tensor)> {
let positions = Tensor::arange(
seqlen_offset as f32,
(seqlen_offset + seq_len) as f32,
self.inv_freq.device(),
)?
.reshape((seq_len, 1))?;
let freqs = positions.matmul(&self.inv_freq)?;
let cos = freqs
.cos()?
.unsqueeze(D::Minus1)?
.repeat((1, 1, 2))?
.flatten_from(D::Minus2)?
.contiguous()?
.to_device(device)?;
let sin = freqs
.sin()?
.unsqueeze(D::Minus1)?
.repeat((1, 1, 2))?
.flatten_from(D::Minus2)?
.contiguous()?
.to_device(device)?;
Ok((cos, sin))
}
}
#[allow(non_camel_case_types)]
#[derive(Debug, Clone)]
pub struct Qwen2_5VLTextRotaryEmbedding {
inv_freq: Vec<f32>,
}
#[allow(non_camel_case_types)]
impl Qwen2_5VLTextRotaryEmbedding {
pub fn new(dim: usize, theta_base: f32) -> Self {
Self {
inv_freq: compute_default_rope_parameters(dim, theta_base),
}
}
pub fn forward(&self, position_ids: &Tensor, dtype: DType, mrope_section: Vec<usize>) -> Result<(Tensor, Tensor)> {
let position_ids_expanded = position_ids.unsqueeze(D::Minus2)?.to_dtype(DType::F32)?.contiguous()?;
let bs = position_ids.dim(1)?;
let inv_freq_expanded = Tensor::from_vec(
self.inv_freq.clone(),
(1, 1, self.inv_freq.len(), 1),
position_ids.device(),
)?
.broadcast_as((3, bs, self.inv_freq.len(), 1))?
.to_dtype(DType::F32)?
.contiguous()?;
let freqs = inv_freq_expanded.matmul(&position_ids_expanded)?.transpose(2, 3)?;
let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?.contiguous()?;
let cos_full = emb.cos()?;
let sin_full = emb.sin()?;
let section_doubled: Vec<usize> = mrope_section.iter().chain(mrope_section.iter()).copied().collect();
let last_dim_full = cos_full.rank() - 1;
let cos_select: Vec<Tensor> = split_tensor(&cos_full, §ion_doubled, last_dim_full)?
.into_iter()
.enumerate()
.map(|(i, m)| m.i(i % 3).map_err(CandleOcrError::from))
.collect::<Result<Vec<_>>>()?;
let cos = Tensor::cat(&cos_select, D::Minus1)?.unsqueeze(1)?.contiguous()?;
let sin_select: Vec<Tensor> = split_tensor(&sin_full, §ion_doubled, last_dim_full)?
.into_iter()
.enumerate()
.map(|(i, m)| m.i(i % 3).map_err(CandleOcrError::from))
.collect::<Result<Vec<_>>>()?;
let sin = Tensor::cat(&sin_select, D::Minus1)?.unsqueeze(1)?.contiguous()?;
Ok((cos.to_dtype(dtype)?, sin.to_dtype(dtype)?))
}
}
#[allow(non_camel_case_types)]
#[derive(Debug, Clone)]
pub struct Qwen2_5VisionRotaryEmbedding {
inv_freq: Vec<f32>,
}
#[allow(non_camel_case_types)]
impl Qwen2_5VisionRotaryEmbedding {
pub fn new(dim: usize, theta_base: Option<f32>) -> Self {
let theta_base = theta_base.unwrap_or(10_000.0_f32);
Self {
inv_freq: compute_default_rope_parameters(dim, theta_base),
}
}
pub fn forward(&self, seqlen: usize, device: &Device) -> Result<Tensor> {
let seq = Tensor::arange(0.0_f32, seqlen as f32, device)?.reshape((seqlen, 1))?;
let inv_freq = Tensor::from_vec(self.inv_freq.clone(), (1, self.inv_freq.len()), device)?;
Ok(seq.matmul(&inv_freq)?)
}
}