#![allow(dead_code)]
use candle_core::{Result, Tensor};
use candle_nn::{LayerNorm, Module, VarBuilder};
use super::config::EncoderConfig;
mod blocks;
mod rope;
use blocks::{ConformerConvolution, ConformerFeedForward, RotaryMHSA, StridingSubsampling};
use rope::create_rope_table;
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)
}
}