2025-09-25 12:09:25 +08:00
|
|
|
use anyhow::Result;
|
2025-10-15 21:03:49 +08:00
|
|
|
use candle_core::{D, Tensor};
|
2025-12-03 17:21:01 +08:00
|
|
|
use candle_nn::{
|
|
|
|
|
Activation, Conv2d, Conv2dConfig, LayerNorm, LayerNormConfig, Linear, Module, VarBuilder,
|
|
|
|
|
conv2d, conv2d_no_bias, layer_norm, linear, linear_no_bias,
|
|
|
|
|
};
|
2025-09-25 12:09:25 +08:00
|
|
|
|
|
|
|
|
use crate::{position_embed::rope::apply_rotary_pos_emb, utils::tensor_utils::repeat_kv};
|
|
|
|
|
|
|
|
|
|
#[derive(Debug, Clone)]
|
2025-12-03 17:21:01 +08:00
|
|
|
pub struct GateUpDownMLP {
|
2025-09-25 12:09:25 +08:00
|
|
|
gate_proj: Linear,
|
|
|
|
|
up_proj: Linear,
|
|
|
|
|
down_proj: Linear,
|
|
|
|
|
act_fn: Activation,
|
|
|
|
|
}
|
|
|
|
|
|
2025-12-03 17:21:01 +08:00
|
|
|
impl GateUpDownMLP {
|
2025-09-25 12:09:25 +08:00
|
|
|
pub fn new(
|
|
|
|
|
vb: VarBuilder,
|
|
|
|
|
hidden_size: usize,
|
|
|
|
|
intermediate_size: usize,
|
|
|
|
|
act_fn: Activation,
|
2025-12-03 17:21:01 +08:00
|
|
|
bias: bool,
|
2025-09-25 12:09:25 +08:00
|
|
|
) -> Result<Self> {
|
2025-12-03 17:21:01 +08:00
|
|
|
let (gate_proj, up_proj, down_proj) = if bias {
|
|
|
|
|
(
|
|
|
|
|
linear(hidden_size, intermediate_size, vb.pp("gate_proj"))?,
|
|
|
|
|
linear(hidden_size, intermediate_size, vb.pp("up_proj"))?,
|
|
|
|
|
linear(intermediate_size, hidden_size, vb.pp("down_proj"))?,
|
|
|
|
|
)
|
|
|
|
|
} else {
|
|
|
|
|
(
|
|
|
|
|
linear_no_bias(hidden_size, intermediate_size, vb.pp("gate_proj"))?,
|
|
|
|
|
linear_no_bias(hidden_size, intermediate_size, vb.pp("up_proj"))?,
|
|
|
|
|
linear_no_bias(intermediate_size, hidden_size, vb.pp("down_proj"))?,
|
|
|
|
|
)
|
|
|
|
|
};
|
2025-09-25 12:09:25 +08:00
|
|
|
Ok(Self {
|
|
|
|
|
gate_proj,
|
|
|
|
|
up_proj,
|
|
|
|
|
down_proj,
|
|
|
|
|
act_fn,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2025-12-03 17:21:01 +08:00
|
|
|
impl Module for GateUpDownMLP {
|
2025-09-25 12:09:25 +08:00
|
|
|
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
|
|
|
|
|
let lhs = xs.apply(&self.gate_proj)?.apply(&self.act_fn)?;
|
|
|
|
|
let rhs = xs.apply(&self.up_proj)?;
|
|
|
|
|
(lhs * rhs)?.apply(&self.down_proj)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2025-12-03 17:21:01 +08:00
|
|
|
pub struct TwoLinearMLP {
|
|
|
|
|
linear1: Linear,
|
|
|
|
|
linear2: Linear,
|
|
|
|
|
act: Activation,
|
2025-09-25 12:09:25 +08:00
|
|
|
}
|
|
|
|
|
|
2025-12-03 17:21:01 +08:00
|
|
|
impl TwoLinearMLP {
|
2025-09-25 12:09:25 +08:00
|
|
|
pub fn new(
|
|
|
|
|
vb: VarBuilder,
|
2025-12-03 17:21:01 +08:00
|
|
|
embedding_dim: usize,
|
|
|
|
|
mlp_dim: usize,
|
|
|
|
|
act: Activation,
|
|
|
|
|
bias: bool,
|
|
|
|
|
linear1_pp_name: &str,
|
|
|
|
|
linear2_pp_name: &str,
|
2025-09-25 12:09:25 +08:00
|
|
|
) -> Result<Self> {
|
2025-12-03 17:21:01 +08:00
|
|
|
let (linear1, linear2) = if bias {
|
|
|
|
|
(
|
|
|
|
|
linear(embedding_dim, mlp_dim, vb.pp(linear1_pp_name))?,
|
|
|
|
|
linear(mlp_dim, embedding_dim, vb.pp(linear2_pp_name))?,
|
|
|
|
|
)
|
|
|
|
|
} else {
|
|
|
|
|
(
|
|
|
|
|
linear_no_bias(embedding_dim, mlp_dim, vb.pp(linear1_pp_name))?,
|
|
|
|
|
linear_no_bias(mlp_dim, embedding_dim, vb.pp(linear2_pp_name))?,
|
|
|
|
|
)
|
|
|
|
|
};
|
2025-09-25 12:09:25 +08:00
|
|
|
Ok(Self {
|
2025-12-03 17:21:01 +08:00
|
|
|
linear1,
|
|
|
|
|
linear2,
|
|
|
|
|
act,
|
2025-09-25 12:09:25 +08:00
|
|
|
})
|
|
|
|
|
}
|
2025-12-03 17:21:01 +08:00
|
|
|
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
|
|
|
|
let xs = xs
|
|
|
|
|
.apply(&self.linear1)?
|
|
|
|
|
.apply(&self.act)?
|
|
|
|
|
.apply(&self.linear2)?;
|
|
|
|
|
Ok(xs)
|
2025-09-25 12:09:25 +08:00
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
#[derive(Debug, Clone)]
|
2025-12-03 17:21:01 +08:00
|
|
|
// pub struct AttentionNobias {
|
|
|
|
|
pub struct NaiveAttention {
|
2025-09-25 12:09:25 +08:00
|
|
|
q_proj: Linear,
|
|
|
|
|
k_proj: Linear,
|
|
|
|
|
v_proj: Linear,
|
|
|
|
|
o_proj: Linear,
|
|
|
|
|
num_heads: usize,
|
|
|
|
|
num_kv_heads: usize,
|
|
|
|
|
num_kv_groups: usize,
|
|
|
|
|
head_dim: usize,
|
2025-12-09 00:41:30 +08:00
|
|
|
middle_size: usize,
|
2025-09-25 12:09:25 +08:00
|
|
|
kv_cache: Option<(Tensor, Tensor)>,
|
|
|
|
|
}
|
|
|
|
|
|
2025-12-03 17:21:01 +08:00
|
|
|
// impl AttentionNobias {
|
|
|
|
|
impl NaiveAttention {
|
2025-10-15 21:03:49 +08:00
|
|
|
pub fn new(
|
|
|
|
|
vb: VarBuilder,
|
|
|
|
|
hidden_size: usize,
|
|
|
|
|
num_attention_heads: usize,
|
|
|
|
|
num_key_value_heads: usize,
|
2025-12-09 00:41:30 +08:00
|
|
|
head_dim: Option<usize>,
|
2025-12-03 17:21:01 +08:00
|
|
|
bias: bool,
|
2025-12-09 00:41:30 +08:00
|
|
|
o_proj_pp_name: Option<&str>,
|
2025-10-15 21:03:49 +08:00
|
|
|
) -> Result<Self> {
|
2025-09-25 12:09:25 +08:00
|
|
|
let num_kv_groups = num_attention_heads / num_key_value_heads;
|
2025-12-09 00:41:30 +08:00
|
|
|
let head_dim = match head_dim {
|
|
|
|
|
None => hidden_size / num_attention_heads,
|
|
|
|
|
Some(dim) => dim,
|
|
|
|
|
};
|
|
|
|
|
let o_proj_pp_name = o_proj_pp_name.unwrap_or("o_proj");
|
2025-12-03 17:21:01 +08:00
|
|
|
let (q_proj, k_proj, v_proj, o_proj) = if bias {
|
|
|
|
|
(
|
|
|
|
|
linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?,
|
|
|
|
|
linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?,
|
|
|
|
|
linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?,
|
2025-12-09 00:41:30 +08:00
|
|
|
linear(
|
|
|
|
|
num_attention_heads * head_dim,
|
|
|
|
|
hidden_size,
|
|
|
|
|
vb.pp(o_proj_pp_name),
|
|
|
|
|
)?,
|
2025-12-03 17:21:01 +08:00
|
|
|
)
|
|
|
|
|
} else {
|
|
|
|
|
(
|
|
|
|
|
linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?,
|
|
|
|
|
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?,
|
|
|
|
|
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?,
|
2025-12-09 00:41:30 +08:00
|
|
|
linear_no_bias(
|
|
|
|
|
num_attention_heads * head_dim,
|
|
|
|
|
hidden_size,
|
|
|
|
|
vb.pp(o_proj_pp_name),
|
|
|
|
|
)?,
|
2025-12-03 17:21:01 +08:00
|
|
|
)
|
|
|
|
|
};
|
|
|
|
|
|
2025-09-25 12:09:25 +08:00
|
|
|
Ok(Self {
|
|
|
|
|
q_proj,
|
|
|
|
|
k_proj,
|
|
|
|
|
v_proj,
|
|
|
|
|
o_proj,
|
|
|
|
|
num_heads: num_attention_heads,
|
|
|
|
|
num_kv_heads: num_key_value_heads,
|
|
|
|
|
num_kv_groups,
|
|
|
|
|
head_dim,
|
2025-12-09 00:41:30 +08:00
|
|
|
middle_size: num_attention_heads * head_dim,
|
2025-09-25 12:09:25 +08:00
|
|
|
kv_cache: None,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
pub fn forward(
|
|
|
|
|
&self,
|
|
|
|
|
xs: &Tensor,
|
2025-12-03 17:21:01 +08:00
|
|
|
cos: Option<&Tensor>,
|
|
|
|
|
sin: Option<&Tensor>,
|
2025-09-25 12:09:25 +08:00
|
|
|
attention_mask: Option<&Tensor>,
|
2025-10-03 22:25:58 +08:00
|
|
|
tof32: bool,
|
2025-09-25 12:09:25 +08:00
|
|
|
) -> Result<Tensor> {
|
|
|
|
|
let (b_sz, q_len, _) = xs.dims3()?;
|
|
|
|
|
let query_states = self.q_proj.forward(xs)?;
|
|
|
|
|
let key_states = self.k_proj.forward(xs)?;
|
|
|
|
|
let value_states = self.v_proj.forward(xs)?;
|
|
|
|
|
let query_states = query_states
|
|
|
|
|
.reshape((b_sz, q_len, self.num_heads, self.head_dim))?
|
|
|
|
|
.transpose(1, 2)?;
|
|
|
|
|
let key_states = key_states
|
|
|
|
|
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
|
|
|
|
|
.transpose(1, 2)?;
|
|
|
|
|
let value_states = value_states
|
|
|
|
|
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
|
|
|
|
|
.transpose(1, 2)?;
|
2025-12-03 17:21:01 +08:00
|
|
|
let (query_states, key_states) = if let Some(cos) = cos
|
|
|
|
|
&& let Some(sin) = sin
|
|
|
|
|
{
|
|
|
|
|
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?
|
|
|
|
|
} else {
|
|
|
|
|
(query_states, key_states)
|
|
|
|
|
};
|
|
|
|
|
|
2025-11-27 18:43:16 +08:00
|
|
|
let scale = 1f64 / f64::sqrt(self.head_dim as f64);
|
|
|
|
|
let attn_output = eager_attention_forward(
|
|
|
|
|
&query_states,
|
|
|
|
|
&key_states,
|
|
|
|
|
&value_states,
|
|
|
|
|
Some(self.num_kv_groups),
|
|
|
|
|
attention_mask,
|
|
|
|
|
scale,
|
|
|
|
|
)?;
|
2025-12-09 00:41:30 +08:00
|
|
|
let attn_output = attn_output.reshape((b_sz, q_len, self.middle_size))?;
|
2025-09-25 12:09:25 +08:00
|
|
|
let attn_output = attn_output.apply(&self.o_proj)?;
|
|
|
|
|
Ok(attn_output)
|
|
|
|
|
}
|
|
|
|
|
|
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>,
|
2025-10-03 22:25:58 +08:00
|
|
|
tof32: bool,
|
2025-09-25 12:09:25 +08:00
|
|
|
) -> Result<Tensor> {
|
|
|
|
|
let (b_sz, q_len, _) = xs.dims3()?;
|
|
|
|
|
let query_states = self.q_proj.forward(xs)?;
|
|
|
|
|
let key_states = self.k_proj.forward(xs)?;
|
|
|
|
|
let value_states = self.v_proj.forward(xs)?;
|
|
|
|
|
let query_states = query_states
|
|
|
|
|
.reshape((b_sz, q_len, self.num_heads, self.head_dim))?
|
|
|
|
|
.transpose(1, 2)?;
|
|
|
|
|
let key_states = key_states
|
|
|
|
|
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
|
|
|
|
|
.transpose(1, 2)?;
|
|
|
|
|
let value_states = value_states
|
|
|
|
|
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
|
|
|
|
|
.transpose(1, 2)?;
|
|
|
|
|
let (query_states, key_states) =
|
2025-10-03 22:25:58 +08:00
|
|
|
apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?;
|
2025-09-25 12:09:25 +08:00
|
|
|
let (key_states, value_states) = match &self.kv_cache {
|
|
|
|
|
None => (key_states, value_states),
|
|
|
|
|
Some((prev_k, prev_v)) => {
|
|
|
|
|
let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
|
|
|
|
|
let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
|
|
|
|
|
(key_states, value_states)
|
|
|
|
|
}
|
|
|
|
|
};
|
|
|
|
|
|
|
|
|
|
self.kv_cache = Some((key_states.clone(), value_states.clone()));
|
2025-11-27 18:43:16 +08:00
|
|
|
let scale = 1f64 / f64::sqrt(self.head_dim as f64);
|
|
|
|
|
let attn_output = eager_attention_forward(
|
|
|
|
|
&query_states,
|
|
|
|
|
&key_states,
|
|
|
|
|
&value_states,
|
|
|
|
|
Some(self.num_kv_groups),
|
|
|
|
|
attention_mask,
|
|
|
|
|
scale,
|
|
|
|
|
)?;
|
2025-12-09 00:41:30 +08:00
|
|
|
let attn_output = attn_output.reshape((b_sz, q_len, self.middle_size))?;
|
2025-09-25 12:09:25 +08:00
|
|
|
let attn_output = attn_output.apply(&self.o_proj)?;
|
|
|
|
|
Ok(attn_output)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
pub fn clear_kv_cache(&mut self) {
|
|
|
|
|
self.kv_cache = None
|
|
|
|
|
}
|
|
|
|
|
}
|
2025-10-26 21:39:23 +08:00
|
|
|
|
|
|
|
|
pub fn eager_attention_forward(
|
|
|
|
|
query_states: &Tensor,
|
|
|
|
|
key_states: &Tensor,
|
|
|
|
|
value_states: &Tensor,
|
|
|
|
|
num_key_value_groups: Option<usize>,
|
|
|
|
|
attention_mask: Option<&Tensor>,
|
|
|
|
|
scaling: f64,
|
|
|
|
|
) -> Result<Tensor> {
|
2025-11-27 18:43:16 +08:00
|
|
|
// input q shape:(b, num_head, seq_len, dim)
|
|
|
|
|
// input k/v shape:(b, num_kv_head, seq_len, dim)
|
2025-10-26 21:39:23 +08:00
|
|
|
let key_states = match num_key_value_groups {
|
|
|
|
|
Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?,
|
2025-10-26 21:57:53 +08:00
|
|
|
None => key_states.clone(),
|
2025-10-26 21:39:23 +08:00
|
|
|
};
|
|
|
|
|
let value_states = match num_key_value_groups {
|
|
|
|
|
Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?,
|
2025-10-26 21:57:53 +08:00
|
|
|
None => value_states.clone(),
|
2025-10-26 21:39:23 +08:00
|
|
|
};
|
2025-11-22 23:27:14 +08:00
|
|
|
let query_states = query_states.contiguous()?;
|
|
|
|
|
let key_states = key_states.contiguous()?;
|
|
|
|
|
let value_states = value_states.contiguous()?;
|
2025-10-26 21:39:23 +08:00
|
|
|
let attn_output = {
|
|
|
|
|
#[cfg(not(feature = "flash-attn"))]
|
|
|
|
|
{
|
|
|
|
|
let attn_weights = query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?)?;
|
2025-10-26 21:57:53 +08:00
|
|
|
let attn_weights = (attn_weights * scaling)?;
|
2025-10-26 21:39:23 +08:00
|
|
|
let attn_weights = match attention_mask {
|
|
|
|
|
None => attn_weights,
|
|
|
|
|
Some(mask) => attn_weights.broadcast_add(&mask.to_dtype(attn_weights.dtype())?)?,
|
|
|
|
|
};
|
|
|
|
|
let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
|
|
|
|
|
attn_weights.matmul(&value_states)?
|
|
|
|
|
}
|
|
|
|
|
#[cfg(feature = "flash-attn")]
|
|
|
|
|
{
|
|
|
|
|
// use flash-attn,
|
|
|
|
|
// flash-attn shape: (bs, seq_len, num_head, head_dim)
|
|
|
|
|
let query_states = query_states.transpose(1, 2)?;
|
|
|
|
|
let key_states = key_states.transpose(1, 2)?;
|
|
|
|
|
let value_states = value_states.transpose(1, 2)?;
|
|
|
|
|
let attn_output = candle_flash_attn::flash_attn(
|
|
|
|
|
&query_states,
|
|
|
|
|
&key_states,
|
|
|
|
|
&value_states,
|
|
|
|
|
scaling as f32,
|
|
|
|
|
attention_mask.is_some(),
|
|
|
|
|
)?
|
|
|
|
|
.transpose(1, 2)?;
|
|
|
|
|
attn_output
|
|
|
|
|
}
|
|
|
|
|
};
|
2025-11-13 00:33:27 +08:00
|
|
|
//(b, n_head, seq_len, dim) -> (b, seq_len, n_head, dim)
|
2025-10-26 21:39:23 +08:00
|
|
|
let attn_output = attn_output.transpose(1, 2)?.contiguous()?;
|
2025-10-26 21:57:53 +08:00
|
|
|
|
2025-10-26 21:39:23 +08:00
|
|
|
Ok(attn_output)
|
|
|
|
|
}
|
2025-12-03 17:21:01 +08:00
|
|
|
|
|
|
|
|
pub fn get_conv2d(
|
|
|
|
|
vb: VarBuilder,
|
|
|
|
|
in_c: usize,
|
|
|
|
|
out_c: usize,
|
|
|
|
|
kernel_size: usize,
|
|
|
|
|
padding: usize,
|
|
|
|
|
stride: usize,
|
|
|
|
|
dilation: usize,
|
|
|
|
|
groups: usize,
|
|
|
|
|
bias: bool,
|
|
|
|
|
) -> Result<Conv2d> {
|
|
|
|
|
let cfg = Conv2dConfig {
|
|
|
|
|
padding,
|
|
|
|
|
stride,
|
|
|
|
|
dilation,
|
|
|
|
|
groups,
|
|
|
|
|
cudnn_fwd_algo: None,
|
|
|
|
|
};
|
|
|
|
|
let conv2d = if bias {
|
|
|
|
|
conv2d(in_c, out_c, kernel_size, cfg, vb)?
|
|
|
|
|
} else {
|
|
|
|
|
conv2d_no_bias(in_c, out_c, kernel_size, cfg, vb)?
|
|
|
|
|
};
|
|
|
|
|
Ok(conv2d)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
pub fn get_layer_norm(vb: VarBuilder, eps: f64, dim: usize) -> Result<LayerNorm> {
|
|
|
|
|
let ln_config = LayerNormConfig {
|
|
|
|
|
eps,
|
|
|
|
|
remove_mean: true, // true for layernorm, false for RMSNorm
|
|
|
|
|
affine: true, // true for with bias, false for without bias
|
|
|
|
|
};
|
|
|
|
|
let norm = layer_norm(dim, ln_config, vb)?;
|
|
|
|
|
Ok(norm)
|
|
|
|
|
}
|