#![allow(dead_code)]
use candle_core::{Device, Result, Tensor};
use candle_nn::{Conv1d, Conv1dConfig, LayerNorm, Linear, Module, VarBuilder};
use super::config::EncoderConfig;
fn create_rope_table(dim: usize, max_len: usize, device: &Device) -> Result<(Tensor, Tensor)> {
let base = 5_000f32;
let half_dim = dim / 2;
let inv_freq: Vec<f32> = (0..half_dim)
.map(|i| 1.0 / base.powf(2.0 * i as f32 / dim as f32))
.collect();
let inv_freq_t = Tensor::from_vec(inv_freq, half_dim, device)?;
let positions: Vec<f32> = (0..max_len).map(|i| i as f32).collect();
let positions_t = Tensor::from_vec(positions, max_len, device)?;
let freqs = positions_t
.unsqueeze(1)?
.matmul(&inv_freq_t.unsqueeze(0)?)?;
let emb = Tensor::cat(&[&freqs, &freqs], 1)?;
let cos = emb.cos()?;
let sin = emb.sin()?;
let cos = cos.unsqueeze(1)?.unsqueeze(1)?;
let sin = sin.unsqueeze(1)?.unsqueeze(1)?;
Ok((cos, sin))
}
fn apply_rotary_pos_emb(
q: &Tensor,
k: &Tensor,
cos: &Tensor,
sin: &Tensor,
) -> Result<(Tensor, Tensor)> {
let q_rot = rotate_half(q)?;
let k_rot = rotate_half(k)?;
let q_embed = q.broadcast_mul(cos)?.add(&q_rot.broadcast_mul(sin)?)?;
let k_embed = k.broadcast_mul(cos)?.add(&k_rot.broadcast_mul(sin)?)?;
Ok((q_embed, k_embed))
}
fn rotate_half(x: &Tensor) -> Result<Tensor> {
let d = x.dim(candle_core::D::Minus1)?;
let half = d / 2;
let x1 = x.narrow(candle_core::D::Minus1, 0, half)?;
let x2 = x.narrow(candle_core::D::Minus1, half, half)?;
let neg_x2 = x2.neg()?;
Tensor::cat(&[&neg_x2, &x1], candle_core::D::Minus1)
}
pub struct StridingSubsampling {
convs: Vec<Conv1d>,
factor: usize,
}
impl StridingSubsampling {
pub fn load(
feat_in: usize,
d_model: usize,
kernel_size: usize,
factor: usize,
vb: VarBuilder,
) -> Result<Self> {
let n_layers = (factor as f64).log2() as usize;
let padding = (kernel_size - 1) / 2;
let cfg = Conv1dConfig {
padding,
stride: 2,
dilation: 1,
groups: 1,
..Default::default()
};
let mut convs = Vec::with_capacity(n_layers);
let mut in_ch = feat_in;
for i in 0..n_layers {
let layer_idx = i * 2;
let conv = candle_nn::conv1d(
in_ch,
d_model,
kernel_size,
cfg,
vb.pp(format!("conv.{layer_idx}")),
)?;
convs.push(conv);
in_ch = d_model;
}
Ok(Self { convs, factor })
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
let mut h = x.transpose(1, 2)?;
for conv in &self.convs {
h = conv.forward(&h)?;
h = h.relu()?;
}
h.transpose(1, 2)
}
pub fn output_length(&self, input_length: usize) -> usize {
let mut length = input_length;
let n_layers = (self.factor as f64).log2() as usize;
for _ in 0..n_layers {
length = length.div_ceil(2);
}
length
}
}
pub struct RotaryMHSA {
linear_q: Linear,
linear_k: Linear,
linear_v: Linear,
linear_out: Linear,
n_heads: usize,
d_k: usize,
}
impl RotaryMHSA {
pub fn load(d_model: usize, n_heads: usize, vb: VarBuilder) -> Result<Self> {
let d_k = d_model / n_heads;
let linear_q = candle_nn::linear(d_model, d_model, vb.pp("linear_q"))?;
let linear_k = candle_nn::linear(d_model, d_model, vb.pp("linear_k"))?;
let linear_v = candle_nn::linear(d_model, d_model, vb.pp("linear_v"))?;
let linear_out = candle_nn::linear(d_model, d_model, vb.pp("linear_out"))?;
Ok(Self {
linear_q,
linear_k,
linear_v,
linear_out,
n_heads,
d_k,
})
}
pub fn forward(
&self,
x: &Tensor,
cos_emb: &Tensor,
sin_emb: &Tensor,
att_mask: Option<&Tensor>,
) -> Result<Tensor> {
let (b, t, _d) = x.dims3()?;
let x_rope = x
.transpose(0, 1)? .reshape((t, b, self.n_heads, self.d_k))?;
let (q_rope, k_rope) = apply_rotary_pos_emb(&x_rope, &x_rope, cos_emb, sin_emb)?;
let q_in = q_rope
.reshape((t, b, self.n_heads * self.d_k))?
.transpose(0, 1)?; let k_in = k_rope
.reshape((t, b, self.n_heads * self.d_k))?
.transpose(0, 1)?; let v_in = x_rope
.reshape((t, b, self.n_heads * self.d_k))?
.transpose(0, 1)?;
let q = self
.linear_q
.forward(&q_in)? .reshape((b, t, self.n_heads, self.d_k))?
.transpose(1, 2)? .contiguous()?;
let k = self
.linear_k
.forward(&k_in)?
.reshape((b, t, self.n_heads, self.d_k))?
.transpose(1, 2)?
.contiguous()?;
let v = self
.linear_v
.forward(&v_in)?
.reshape((b, t, self.n_heads, self.d_k))?
.transpose(1, 2)?
.contiguous()?;
let scale = (self.d_k as f64).sqrt();
let mut scores = q.matmul(&k.transpose(2, 3)?)?;
scores = (scores / scale)?;
if let Some(mask) = att_mask {
let mask = mask.unsqueeze(1)?;
let fill_val =
Tensor::new(-10_000f32, scores.device())?.broadcast_as(scores.shape())?;
scores = mask.where_cond(&fill_val, &scores)?;
}
let attn = candle_nn::ops::softmax_last_dim(&scores)?;
if let Some(mask) = att_mask {
let mask = mask.unsqueeze(1)?;
let zeros = Tensor::zeros_like(&attn)?;
let attn = mask.where_cond(&zeros, &attn)?;
let context = attn.matmul(&v)?; let context = context
.transpose(1, 2)? .reshape((b, t, self.n_heads * self.d_k))?;
return self.linear_out.forward(&context);
}
let context = attn.matmul(&v)?;
let context = context
.transpose(1, 2)?
.reshape((b, t, self.n_heads * self.d_k))?;
self.linear_out.forward(&context)
}
}
pub struct ConformerFeedForward {
linear1: Linear,
linear2: Linear,
}
impl ConformerFeedForward {
pub fn load(d_model: usize, d_ff: usize, vb: VarBuilder) -> Result<Self> {
let linear1 = candle_nn::linear(d_model, d_ff, vb.pp("linear1"))?;
let linear2 = candle_nn::linear(d_ff, d_model, vb.pp("linear2"))?;
Ok(Self { linear1, linear2 })
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
let h = self.linear1.forward(x)?;
let h = candle_nn::Activation::Silu.forward(&h)?;
self.linear2.forward(&h)
}
}
pub struct ConformerConvolution {
pointwise_conv1: Conv1d,
depthwise_conv: Conv1d,
norm: LayerNorm,
pointwise_conv2: Conv1d,
d_model: usize,
}
impl ConformerConvolution {
pub fn load(d_model: usize, kernel_size: usize, vb: VarBuilder) -> Result<Self> {
let padding = (kernel_size - 1) / 2;
let pw1_cfg = Conv1dConfig {
padding: 0,
stride: 1,
dilation: 1,
groups: 1,
..Default::default()
};
let pointwise_conv1 =
candle_nn::conv1d(d_model, d_model * 2, 1, pw1_cfg, vb.pp("pointwise_conv1"))?;
let dw_cfg = Conv1dConfig {
padding,
stride: 1,
dilation: 1,
groups: d_model,
..Default::default()
};
let depthwise_conv = candle_nn::conv1d(
d_model,
d_model,
kernel_size,
dw_cfg,
vb.pp("depthwise_conv"),
)?;
let norm = candle_nn::layer_norm(d_model, 1e-5, vb.pp("batch_norm"))?;
let pw2_cfg = Conv1dConfig {
padding: 0,
stride: 1,
dilation: 1,
groups: 1,
..Default::default()
};
let pointwise_conv2 =
candle_nn::conv1d(d_model, d_model, 1, pw2_cfg, vb.pp("pointwise_conv2"))?;
Ok(Self {
pointwise_conv1,
depthwise_conv,
norm,
pointwise_conv2,
d_model,
})
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
let h = x.transpose(1, 2)?;
let h = self.pointwise_conv1.forward(&h)?;
let h1 = h.narrow(1, 0, self.d_model)?;
let h2 = h.narrow(1, self.d_model, self.d_model)?;
let h = (h1 * candle_nn::ops::sigmoid(&h2)?)?;
let h = self.depthwise_conv.forward(&h)?;
let h = h.transpose(1, 2)?; let h = self.norm.forward(&h)?;
let h = h.transpose(1, 2)?;
let h = candle_nn::Activation::Silu.forward(&h)?;
let h = self.pointwise_conv2.forward(&h)?;
h.transpose(1, 2)
}
}
pub struct ConformerLayer {
norm_feed_forward1: LayerNorm,
feed_forward1: ConformerFeedForward,
norm_self_att: LayerNorm,
self_attn: RotaryMHSA,
norm_conv: LayerNorm,
conv: ConformerConvolution,
norm_feed_forward2: LayerNorm,
feed_forward2: ConformerFeedForward,
norm_out: LayerNorm,
}
impl ConformerLayer {
pub fn load(
d_model: usize,
d_ff: usize,
n_heads: usize,
conv_kernel_size: usize,
vb: VarBuilder,
) -> Result<Self> {
let norm_feed_forward1 = candle_nn::layer_norm(d_model, 1e-5, vb.pp("norm_feed_forward1"))?;
let feed_forward1 = ConformerFeedForward::load(d_model, d_ff, vb.pp("feed_forward1"))?;
let norm_self_att = candle_nn::layer_norm(d_model, 1e-5, vb.pp("norm_self_att"))?;
let self_attn = RotaryMHSA::load(d_model, n_heads, vb.pp("self_attn"))?;
let norm_conv = candle_nn::layer_norm(d_model, 1e-5, vb.pp("norm_conv"))?;
let conv = ConformerConvolution::load(d_model, conv_kernel_size, vb.pp("conv"))?;
let norm_feed_forward2 = candle_nn::layer_norm(d_model, 1e-5, vb.pp("norm_feed_forward2"))?;
let feed_forward2 = ConformerFeedForward::load(d_model, d_ff, vb.pp("feed_forward2"))?;
let norm_out = candle_nn::layer_norm(d_model, 1e-5, vb.pp("norm_out"))?;
Ok(Self {
norm_feed_forward1,
feed_forward1,
norm_self_att,
self_attn,
norm_conv,
conv,
norm_feed_forward2,
feed_forward2,
norm_out,
})
}
pub fn forward(
&self,
x: &Tensor,
cos_emb: &Tensor,
sin_emb: &Tensor,
att_mask: Option<&Tensor>,
) -> Result<Tensor> {
const FC_FACTOR: f64 = 0.5;
let h = self.norm_feed_forward1.forward(x)?;
let h = self.feed_forward1.forward(&h)?;
let residual = (x + (h * FC_FACTOR)?)?;
let h = self.norm_self_att.forward(&residual)?;
let h = self.self_attn.forward(&h, cos_emb, sin_emb, att_mask)?;
let residual = (residual + h)?;
let h = self.norm_conv.forward(&residual)?;
let h = self.conv.forward(&h)?;
let residual = (residual + h)?;
let h = self.norm_feed_forward2.forward(&residual)?;
let h = self.feed_forward2.forward(&h)?;
let residual = (residual + (h * FC_FACTOR)?)?;
self.norm_out.forward(&residual)
}
}
pub struct ConformerEncoder {
pre_encode: StridingSubsampling,
layers: Vec<ConformerLayer>,
rope_cos: Tensor,
rope_sin: Tensor,
d_k: usize,
#[allow(dead_code)]
feat_in: usize,
}
impl ConformerEncoder {
pub fn load(config: &EncoderConfig, vb: VarBuilder) -> Result<Self> {
let d_k = config.d_model / config.n_heads;
let d_ff = config.d_model * config.ff_expansion_factor;
let pre_encode = StridingSubsampling::load(
config.feat_in,
config.d_model,
config.subs_kernel_size,
config.subsampling_factor,
vb.pp("pre_encode"),
)?;
let (rope_cos, rope_sin) = create_rope_table(d_k, config.pos_emb_max_len, vb.device())?;
let rope_cos = rope_cos.to_dtype(vb.dtype())?;
let rope_sin = rope_sin.to_dtype(vb.dtype())?;
let mut layers = Vec::with_capacity(config.n_layers);
for i in 0..config.n_layers {
let layer = ConformerLayer::load(
config.d_model,
d_ff,
config.n_heads,
config.conv_kernel_size,
vb.pp(format!("layers.{i}")),
)?;
layers.push(layer);
}
Ok(Self {
pre_encode,
layers,
rope_cos,
rope_sin,
d_k,
feat_in: config.feat_in,
})
}
pub fn forward(&self, features: &Tensor) -> Result<Tensor> {
let x = features.transpose(1, 2)?;
let x = self.pre_encode.forward(&x)?;
let (_b, t, _d) = x.dims3()?;
let (cos_emb, sin_emb) = if t <= self.rope_cos.dim(0)? {
(
self.rope_cos.narrow(0, 0, t)?,
self.rope_sin.narrow(0, 0, t)?,
)
} else {
tracing::warn!(
"GigaAM: RoPE таблица расширена с {} до {} позиций",
self.rope_cos.dim(0)?,
t,
);
let (cos, sin) = create_rope_table(self.d_k, t, x.device())?;
(cos.to_dtype(x.dtype())?, sin.to_dtype(x.dtype())?)
};
const SYNC_EVERY: usize = 4;
let is_metal = x.device().is_metal();
let mut h = x;
for (i, layer) in self.layers.iter().enumerate() {
h = layer.forward(&h, &cos_emb, &sin_emb, None)?;
if is_metal && (i + 1) % SYNC_EVERY == 0 {
h.device().synchronize().map_err(|e| {
candle_core::Error::Msg(format!("Metal sync at layer {}: {e}", i + 1))
})?;
}
}
h.transpose(1, 2)
}
}