use candle_core::{Result, Tensor};
use candle_nn::{Conv1d, Conv1dConfig, LayerNorm, Linear, Module, VarBuilder};
use super::rope::apply_rotary_pos_emb;
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)
}
}