Files
aha/src/models/minicpm4/model.rs
T

338 lines
11 KiB
Rust
Raw Normal View History

2025-10-15 21:03:49 +08:00
use anyhow::{Ok, Result};
use candle_core::{D, Device, Tensor};
use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_norm};
2025-09-25 12:09:25 +08:00
use crate::{
models::{
2026-04-02 22:22:52 +08:00
common::{
InferenceModel,
modules::{GateUpDownMLP, NaiveAttention},
},
2025-09-25 12:09:25 +08:00
minicpm4::config::MiniCPM4Config,
},
position_embed::rope::compute_default_rope_parameters,
utils::tensor_utils::prepare_causal_attention_mask,
};
pub struct MiniCPMLongRoPE {
short_factor: Vec<f32>,
long_factor: Vec<f32>,
original_max_position_embeddings: usize,
2025-10-03 22:25:58 +08:00
max_seq_len_cached: usize,
scaling_factor: f64,
2025-09-25 12:09:25 +08:00
inv_freq: Tensor,
cos_cached: Tensor,
sin_cached: Tensor,
2025-10-03 22:25:58 +08:00
device: Device,
2025-09-25 12:09:25 +08:00
}
impl MiniCPMLongRoPE {
pub fn new(cfg: &MiniCPM4Config, device: &Device) -> Result<Self> {
let head_dim = cfg.hidden_size / cfg.num_attention_heads;
let rope_theta = 10000.0;
let short_factor = cfg.rope_scaling.short_factor.clone();
let long_factor = cfg.rope_scaling.short_factor.clone();
let original_max_position_embeddings = cfg.rope_scaling.original_max_position_embeddings;
2025-10-03 22:25:58 +08:00
let max_position_embeddings = cfg.max_position_embeddings;
let scale = max_position_embeddings as f64 / original_max_position_embeddings as f64;
2025-09-25 12:09:25 +08:00
let scaling_factor =
2025-10-03 22:25:58 +08:00
(1.0 + scale.ln() / (original_max_position_embeddings as f64).ln()).sqrt();
2025-09-25 12:09:25 +08:00
let inv_freq = compute_default_rope_parameters(head_dim, rope_theta);
2025-10-15 21:03:49 +08:00
let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?;
2025-10-03 22:25:58 +08:00
let max_seq_len_cached = max_position_embeddings;
2025-09-25 12:09:25 +08:00
let t = Tensor::arange(0.0_f32, max_position_embeddings as f32, device)?
.reshape((max_position_embeddings, 1))?;
// short_factor.len() = 32
// head_dim = 1024 / 16 = 64, inv_freq.len() = 32
2025-10-15 21:03:49 +08:00
let ext_factors = Tensor::from_slice(&short_factor, (1, short_factor.len()), device)?;
2025-09-25 12:09:25 +08:00
let ext_factors = Tensor::ones_like(&ext_factors)?.div(&ext_factors)?;
// (seq_len, 1) matmul (1, 32) -> (seq_len, 32) * (1, 32)-> (seq_len, 32)
let freqs = t.matmul(&ext_factors)?.broadcast_mul(&inv_freq)?;
let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?;
let cos_cached = emb.cos()?.affine(scaling_factor, 0.0)?;
let sin_cached = emb.sin()?.affine(scaling_factor, 0.0)?;
Ok(Self {
short_factor,
long_factor,
original_max_position_embeddings,
2025-10-03 22:25:58 +08:00
max_seq_len_cached,
scaling_factor,
2025-09-25 12:09:25 +08:00
inv_freq,
cos_cached,
sin_cached,
2025-10-03 22:25:58 +08:00
device: device.clone(),
2025-09-25 12:09:25 +08:00
})
}
2025-10-03 22:25:58 +08:00
pub fn update_cos_sin_cache(&mut self, seqlen: usize) -> Result<()> {
self.max_seq_len_cached = seqlen;
2025-10-15 21:03:49 +08:00
let t = Tensor::arange(0.0_f32, seqlen as f32, &self.device)?.reshape((seqlen, 1))?;
2025-10-03 22:25:58 +08:00
let mut ext_factors = Tensor::from_slice(
&self.short_factor,
(1, self.short_factor.len()),
&self.device,
)?;
2025-09-25 12:09:25 +08:00
if seqlen > self.original_max_position_embeddings {
ext_factors =
2025-10-03 22:25:58 +08:00
Tensor::from_slice(&self.long_factor, (1, self.long_factor.len()), &self.device)?;
2025-09-25 12:09:25 +08:00
}
let ext_factors = Tensor::ones_like(&ext_factors)?.div(&ext_factors)?;
let freqs = t.matmul(&ext_factors)?.broadcast_mul(&self.inv_freq)?;
let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?;
2025-10-03 22:25:58 +08:00
let cos_cached = emb.cos()?.affine(self.scaling_factor, 0.0)?;
let sin_cached = emb.sin()?.affine(self.scaling_factor, 0.0)?;
2025-09-25 12:09:25 +08:00
self.cos_cached = cos_cached;
self.sin_cached = sin_cached;
Ok(())
}
2025-10-03 22:25:58 +08:00
pub fn forward(&mut self, pos_offset: usize, seqlen: usize) -> Result<(Tensor, Tensor)> {
if pos_offset + seqlen > self.max_seq_len_cached {
2025-10-15 21:03:49 +08:00
self.update_cos_sin_cache(pos_offset + seqlen)?;
2025-10-03 22:25:58 +08:00
}
2025-09-25 12:09:25 +08:00
let cos = self.cos_cached.narrow(0, pos_offset, seqlen)?;
let sin = self.sin_cached.narrow(0, pos_offset, seqlen)?;
2025-10-03 22:25:58 +08:00
2025-09-25 12:09:25 +08:00
Ok((cos, sin))
}
}
pub struct MiniCPMDecoderLayer {
2025-12-03 17:21:01 +08:00
self_attn: NaiveAttention,
mlp: GateUpDownMLP,
2025-09-25 12:09:25 +08:00
input_layernorm: RmsNorm,
post_attention_layernorm: RmsNorm,
scale_depth: f32,
num_hidden_layers: usize,
}
impl MiniCPMDecoderLayer {
pub fn new(vb: VarBuilder, cfg: &MiniCPM4Config) -> Result<Self> {
2025-12-03 17:21:01 +08:00
let self_attn = NaiveAttention::new(
2025-09-25 12:09:25 +08:00
vb.pp("self_attn"),
cfg.hidden_size,
cfg.num_attention_heads,
cfg.num_key_value_heads,
2025-12-09 00:41:30 +08:00
None,
2025-12-03 17:21:01 +08:00
false,
2025-12-09 00:41:30 +08:00
None,
2026-01-15 21:57:12 +08:00
None,
None,
None,
2025-09-25 12:09:25 +08:00
)?;
2025-12-03 17:21:01 +08:00
let mlp = GateUpDownMLP::new(
2025-09-25 12:09:25 +08:00
vb.pp("mlp"),
cfg.hidden_size,
cfg.intermediate_size,
cfg.hidden_act,
2025-12-03 17:21:01 +08:00
false,
2026-01-30 22:04:23 +08:00
None,
None,
None,
2025-09-25 12:09:25 +08:00
)?;
let input_layernorm =
rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
let post_attention_layernorm = rms_norm(
cfg.hidden_size,
cfg.rms_norm_eps,
vb.pp("post_attention_layernorm"),
)?;
Ok(Self {
self_attn,
mlp,
input_layernorm,
post_attention_layernorm,
scale_depth: cfg.scale_depth,
num_hidden_layers: cfg.num_hidden_layers,
})
}
pub fn forward(
&self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
2025-10-03 22:25:58 +08:00
let residual = xs.clone();
2025-09-25 12:09:25 +08:00
let xs = self.input_layernorm.forward(xs)?;
2025-10-15 21:03:49 +08:00
let xs = self
.self_attn
2025-12-03 17:21:01 +08:00
.forward(&xs, Some(cos), Some(sin), attention_mask, true)?;
2025-10-03 22:25:58 +08:00
let xs = (residual
+ xs.affine(
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
0.0,
))?;
2025-09-25 12:09:25 +08:00
let residual = &xs;
let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?;
2025-10-03 22:25:58 +08:00
let xs = (residual
+ xs.affine(
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
0.0,
2025-10-10 23:25:25 +08:00
)?)?;
2025-09-25 12:09:25 +08:00
Ok(xs)
}
2025-10-10 20:36:52 +08:00
pub fn forward_with_cache(
2025-09-25 12:09:25 +08:00
&mut self,
xs: &Tensor,
cos: &Tensor,
sin: &Tensor,
attention_mask: Option<&Tensor>,
) -> Result<Tensor> {
2025-10-03 22:25:58 +08:00
let residual = xs.clone();
2025-09-25 12:09:25 +08:00
let xs = self.input_layernorm.forward(xs)?;
2026-04-08 18:55:03 +08:00
let xs =
self.self_attn
.forward_with_cache(&xs, Some(cos), Some(sin), attention_mask, true)?;
2025-10-03 22:25:58 +08:00
let xs = (residual
+ xs.affine(
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
0.0,
))?;
2025-09-25 12:09:25 +08:00
let residual = &xs;
let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?;
2025-10-03 22:25:58 +08:00
let xs = (residual
+ xs.affine(
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
0.0,
2025-10-10 23:25:25 +08:00
)?)?;
2025-09-25 12:09:25 +08:00
Ok(xs)
}
pub fn clear_kv_cache(&mut self) {
self.self_attn.clear_kv_cache();
}
}
pub struct MiniCPMModel {
cfg: MiniCPM4Config,
embed_tokens: Embedding,
layers: Vec<MiniCPMDecoderLayer>,
norm: RmsNorm,
rope_emb: MiniCPMLongRoPE,
lm_head: Linear,
2026-04-02 22:22:52 +08:00
stop_token_ids: Vec<u32>,
2025-09-25 12:09:25 +08:00
}
impl MiniCPMModel {
2026-05-27 00:50:23 +08:00
pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> Result<Self> {
2025-10-03 22:25:58 +08:00
let vb = vb.pp("model");
2025-09-25 12:09:25 +08:00
let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?;
let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
let vb_layers = vb.pp("layers");
for i in 0..cfg.num_hidden_layers {
let layer = MiniCPMDecoderLayer::new(vb_layers.pp(i), &cfg)?;
layers.push(layer);
}
let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?;
let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device())?;
let lm_head = Linear::new(embed_tokens.embeddings().clone(), None);
2026-05-27 00:50:23 +08:00
let stop_token_ids = cfg.eos_token_id.clone();
2025-09-25 12:09:25 +08:00
Ok(Self {
cfg,
embed_tokens,
layers,
norm,
rope_emb,
2025-10-03 22:25:58 +08:00
lm_head,
2026-05-27 00:50:23 +08:00
stop_token_ids,
2025-09-25 12:09:25 +08:00
})
}
2026-04-02 22:22:52 +08:00
pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
2025-09-25 12:09:25 +08:00
let (bs, seq_len) = input_ids.dims2()?;
2025-10-03 22:25:58 +08:00
let input_embeds = self
.embed_tokens
2025-10-15 21:03:49 +08:00
.forward(input_ids)?
2025-10-03 22:25:58 +08:00
.affine(self.cfg.scale_emb, 0.0)?;
let attention_mask: Option<Tensor> = {
2025-09-25 12:09:25 +08:00
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
2025-09-25 12:09:25 +08:00
bs,
seq_len,
2025-10-11 13:14:51 +08:00
0,
2025-09-25 12:09:25 +08:00
input_ids.device(),
)?)
}
};
2025-10-15 21:03:49 +08:00
2026-04-02 22:22:52 +08:00
let (cos, sin) = self.rope_emb.forward(seqlen_offset, seq_len)?;
2025-09-25 12:09:25 +08:00
let mut hidden_states = input_embeds;
for decode_layer in &self.layers {
hidden_states =
decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?;
2025-09-25 12:09:25 +08:00
}
hidden_states = self.norm.forward(&hidden_states)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
2025-10-03 22:25:58 +08:00
let hidden_state = hidden_state.affine(
1.0 / (self.cfg.hidden_size / self.cfg.dim_model_base) as f64,
0.0,
)?;
2025-09-25 12:09:25 +08:00
let logits = self.lm_head.forward(&hidden_state)?;
Ok(logits)
}
2026-04-02 22:22:52 +08:00
pub fn forward_with_cache(
&mut self,
input_ids: &Tensor,
seqlen_offset: usize,
) -> Result<Tensor> {
2025-09-25 12:09:25 +08:00
let (bs, seq_len) = input_ids.dims2()?;
2025-10-03 22:25:58 +08:00
let input_embeds = self
.embed_tokens
2025-10-15 21:03:49 +08:00
.forward(input_ids)?
2025-10-03 22:25:58 +08:00
.affine(self.cfg.scale_emb, 0.0)?;
let attention_mask: Option<Tensor> = {
2025-09-25 12:09:25 +08:00
if seq_len <= 1 {
None
} else {
Some(prepare_causal_attention_mask(
2025-09-25 12:09:25 +08:00
bs,
seq_len,
2025-10-11 13:14:51 +08:00
0,
2025-09-25 12:09:25 +08:00
input_ids.device(),
)?)
}
};
2026-04-02 22:22:52 +08:00
let (cos, sin) = self.rope_emb.forward(seqlen_offset, seq_len)?;
2025-09-25 12:09:25 +08:00
let mut hidden_states = input_embeds;
for decode_layer in &mut self.layers {
hidden_states = decode_layer.forward_with_cache(
&hidden_states,
&cos,
&sin,
attention_mask.as_ref(),
)?;
2025-09-25 12:09:25 +08:00
}
hidden_states = self.norm.forward(&hidden_states)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
2025-10-03 22:25:58 +08:00
let hidden_state = hidden_state.affine(
1.0 / (self.cfg.hidden_size / self.cfg.dim_model_base) as f64,
0.0,
)?;
2025-09-25 12:09:25 +08:00
let logits = self.lm_head.forward(&hidden_state)?;
Ok(logits)
}
pub fn clear_kv_cache(&mut self) {
for layer in self.layers.iter_mut() {
layer.clear_kv_cache()
}
}
}
2026-04-02 22:22:52 +08:00
impl InferenceModel for MiniCPMModel {
fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
self.forward_with_cache(input_ids, seqlen_offset)
}
fn clear_cache(&mut self) {
self.clear_kv_cache();
}
fn stop_token_ids(&self) -> Vec<u32> {
self.stop_token_ids.clone()
}
}