stash save

This commit is contained in:
jhqxxx
2026-05-09 12:18:16 +08:00
parent 5df1f30f83
commit e8bf980b38
39 changed files with 2505 additions and 316 deletions
+118 -56
View File
@@ -1,4 +1,4 @@
use anyhow::{Result, anyhow};
use anyhow::Result;
use candle_core::{D, DType, Device, IndexOp, Tensor};
use candle_transformers::models::deepseek2::SplitOp;
@@ -21,6 +21,22 @@ pub fn rotate_half(x: &Tensor) -> Result<Tensor> {
Ok(rotate_x)
}
pub fn rotate_half_interleave(x: &Tensor) -> Result<Tensor> {
let x_rank = x.rank();
let x_dim = x.dims();
let half_dim = x_dim[x_rank - 1] / 2;
let mut x_reshape = x_dim[0..x_rank - 1].to_vec();
x_reshape.push(half_dim);
x_reshape.push(2);
let x = x.reshape(x_reshape)?;
let even = x.narrow(D::Minus1, 0, 1)?;
let odd = x.narrow(D::Minus1, 1, 1)?.affine(-1.0, 0.0)?;
let rotate_x = Tensor::cat(&[&odd, &even], D::Minus1)?
.reshape(x_dim)?
.contiguous()?;
Ok(rotate_x)
}
pub fn apply_multimodel_rotary_pos_emb(
q: &Tensor,
k: &Tensor,
@@ -115,6 +131,44 @@ pub fn apply_rotary_pos_emb(
Ok((q_embed, k_embed))
}
pub fn apply_rotary_pos_emb_interleave(
q: &Tensor,
k: &Tensor,
cos: &Tensor,
sin: &Tensor,
tof32: bool,
) -> Result<(Tensor, Tensor)> {
// sin/cos: to (bs, 1, seq_len, head_dim)
// q/k: (bs, n_head, seq_len, head_dim)
let mut cos = cos.clone();
let mut sin = sin.clone();
if cos.rank() == 2 {
// (seq_len, head_dim) -> (1, 1, seq_len, head_dim)
cos = cos.unsqueeze(0)?.unsqueeze(0)?;
sin = sin.unsqueeze(0)?.unsqueeze(0)?;
}
if cos.rank() == 3 {
// (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim)
cos = cos.unsqueeze(1)?;
sin = sin.unsqueeze(1)?;
}
let orig_dtype = q.dtype();
let q = if tof32 { &q.to_dtype(DType::F32)? } else { q };
let k = if tof32 { &k.to_dtype(DType::F32)? } else { k };
let cos = cos.to_dtype(q.dtype())?;
let sin = sin.to_dtype(q.dtype())?;
let q_embed = q
.broadcast_mul(&cos)?
.add(&rotate_half_interleave(q)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
let k_embed = k
.broadcast_mul(&cos)?
.add(&rotate_half_interleave(k)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
Ok((q_embed, k_embed))
}
pub fn glm_asr_apply_rotary_pos_emb(
q: &Tensor,
k: &Tensor,
@@ -258,66 +312,46 @@ pub fn glm_ocr_apply_rotary_pos_emb(
Ok((q_embed, k_embed))
}
pub fn roformer_rotate(x: &Tensor) -> Result<Tensor> {
let dims = x.dims();
let last_dim = dims
.last()
.ok_or(anyhow!("Input tensor must have at least one dimension"))?;
if last_dim % 2 != 0 {
return Err(anyhow!(
"Last dimension size must be even, got {}",
last_dim
));
}
let new_dims: Vec<usize> = dims[..dims.len() - 1]
.iter()
.copied()
.chain([last_dim / 2, 2])
.collect();
let x_reshape = x.reshape(new_dims)?;
let x_chunks = x_reshape.chunk(2, D::Minus1)?;
let x1 = &x_chunks[0];
let x2 = &x_chunks[1];
// let x1 = x_reshape.narrow(D::Minus1, 0, 1)?;
// let x2 = x_reshape.narrow(D::Minus1, 1, 1)?;
let x2_neg = x2.affine(-1.0, 0.0)?;
let rotate_x = Tensor::cat(&[&x2_neg, x1], D::Minus1)?;
Ok(rotate_x.flatten(D::Minus2, D::Minus1)?)
}
pub fn apply_rotary_pos_emb_roformer(
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 {
// (seq_len, head_dim) -> (1, 1, seq_len, head_dim)
cos = cos.unsqueeze(0)?.unsqueeze(0)?;
sin = sin.unsqueeze(0)?.unsqueeze(0)?;
}
if cos.rank() == 3 {
// (bs, seq_len, head_dim) -> (bs, 1, seq_len, head_dim)
cos = cos.unsqueeze(1)?;
sin = sin.unsqueeze(1)?;
}
let orig_dtype = q.dtype();
let q = if tof32 { &q.to_dtype(DType::F32)? } else { q };
let k = if tof32 { &k.to_dtype(DType::F32)? } else { k };
let cos = cos.to_dtype(q.dtype())?;
let sin = sin.to_dtype(q.dtype())?;
let q_embed = q
.broadcast_mul(&cos)?
.add(&roformer_rotate(q)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
let k_embed = k
.broadcast_mul(&cos)?
.add(&roformer_rotate(k)?.broadcast_mul(&sin)?)?
.to_dtype(orig_dtype)?;
Ok((q_embed, k_embed))
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(candle_core::DType::F32)?;
let roti = sin
.narrow(D::Minus1, 0, half_dim)?
.to_dtype(candle_core::DType::F32)?;
let q = q
.reshape((bs, n_head, seq_len, half_dim, 2))?
.to_dtype(candle_core::DType::F32)?;
let qr = q.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let qi = q.narrow(D::Minus1, 1, 1)?.squeeze(D::Minus1)?;
let k = k
.reshape((bs, n_head, seq_len, half_dim, 2))?
.to_dtype(candle_core::DType::F32)?;
let kr = k.narrow(D::Minus1, 0, 1)?.squeeze(D::Minus1)?;
let ki = k.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 = Tensor::stack(&[qor, qoi], D::Minus1)?
.reshape((bs, n_head, seq_len, dim))?
.to_dtype(ori_dtype)?;
let k = Tensor::stack(&[kor, koi], D::Minus1)?
.reshape((bs, n_head, seq_len, dim))?
.to_dtype(ori_dtype)?;
Ok((q, k))
}
#[derive(Debug, Clone)]
@@ -554,7 +588,6 @@ 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(
@@ -577,6 +610,35 @@ impl RoPE {
let sin = emb.sin()?;
Ok((cos, 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))?; // (seq_len, 1)
let freqs = positions.matmul(&self.inv_freq)?; // (seq_len, dim / 2)
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))
}
}
pub fn get_xd_cos_sin(