use burn::{
config::Config,
module::Module,
nn::{
LinearLayout,
PaddingConfig1d,
activation::ActivationConfig,
conv::{
Conv1d,
Conv1dConfig,
},
},
prelude::{
Backend,
Tensor,
s,
},
tensor::{
activation::{
relu,
sigmoid,
},
ops::PadMode,
},
};
use crate::{
blocks::conv::{
ConvBlock1dConfig,
ConvBlock1dMeta,
ConvSeq1d,
ConvSeq1dConfig,
ConvSeq1dMeta,
},
burner::module::ModuleInit,
errors::{
BunsenError,
BunsenResult,
},
kits::speech::silero_vad::{
FusedLstm,
FusedLstmConfig,
blocks::context::SileroVadContext,
},
};
#[derive(Config, Debug)]
pub struct SileroVadSignalConfig {
pub sample_rate: usize,
pub n_freq: usize,
#[config(default = "128")]
pub d_hidden: usize,
#[config(default = "64")]
pub d_bottleneck: usize,
}
impl SileroVadSignalConfig {
pub fn standard_16khz() -> Self {
Self::new(16000, 129)
}
pub fn standard_8khz() -> Self {
Self::new(8000, 65)
}
pub fn to_stft(&self) -> SileroVadStftConfig {
let stft_stride = self.n_freq - 1;
let stft_kernel = stft_stride * 2;
let input_pad = stft_stride / 2;
SileroVadStftConfig::new(
self.sample_rate,
self.n_freq,
input_pad,
stft_kernel,
stft_stride,
)
.with_d_hidden(self.d_hidden)
.with_d_bottleneck(self.d_bottleneck)
}
pub fn to_structure(&self) -> SileroVadStructureConfig {
self.to_stft().to_structure()
}
}
impl<B: Backend> ModuleInit<B, SileroVad<B>> for SileroVadSignalConfig {
fn try_init(
&self,
device: &B::Device,
) -> BunsenResult<SileroVad<B>> {
self.to_stft().try_init(device)
}
}
#[derive(Config, Debug)]
pub struct SileroVadStftConfig {
pub sample_rate: usize,
pub n_freq: usize,
pub input_pad: usize,
pub stft_kernel: usize,
pub stft_stride: usize,
#[config(default = "128")]
pub d_hidden: usize,
#[config(default = "64")]
pub d_bottleneck: usize,
}
impl SileroVadStftConfig {
pub fn to_structure(&self) -> SileroVadStructureConfig {
SileroVadStructureConfig {
sample_rate: self.sample_rate,
input_pad: self.input_pad,
stft: Conv1dConfig::new(1, 2 * self.n_freq, self.stft_kernel)
.with_stride(self.stft_stride)
.with_padding(PaddingConfig1d::Valid)
.with_bias(false),
encoder: encoder_config(self.n_freq, self.d_hidden, self.d_bottleneck),
lstm: FusedLstmConfig::new(self.d_hidden).with_layout(LinearLayout::Col),
decoder: Conv1dConfig::new(self.d_hidden, 1, 1)
.with_padding(PaddingConfig1d::Valid)
.with_bias(true),
}
}
}
impl<B: Backend> ModuleInit<B, SileroVad<B>> for SileroVadStftConfig {
fn try_init(
&self,
device: &B::Device,
) -> BunsenResult<SileroVad<B>> {
self.to_structure().try_init(device)
}
}
pub trait SileroVadMeta {
fn sample_rate(&self) -> usize;
fn n_freq(&self) -> usize;
fn chunk_size(&self) -> usize {
2 * self.stft_kernel()
}
fn input_pad(&self) -> usize;
fn stft_kernel(&self) -> usize;
fn stft_stride(&self) -> usize;
fn d_hidden(&self) -> usize;
fn d_bottleneck(&self) -> usize;
fn gate_size(&self) -> usize {
4 * self.d_hidden()
}
}
pub fn encoder_config(
n_freq: usize,
d_hidden: usize,
d_bottleneck: usize,
) -> ConvSeq1dConfig {
let block = |in_channels: usize, out_channels: usize, stride: usize| {
ConvBlock1dConfig::new(
Conv1dConfig::new(in_channels, out_channels, 3)
.with_stride(stride)
.with_padding(PaddingConfig1d::Explicit(1, 1))
.with_bias(true),
)
.with_act(Some(ActivationConfig::Relu))
};
ConvSeq1dConfig::new(vec![
block(n_freq, d_hidden, 1),
block(d_hidden, d_bottleneck, 2),
block(d_bottleneck, d_bottleneck, 2),
block(d_bottleneck, d_hidden, 1),
])
}
#[derive(Config, Debug)]
pub struct SileroVadStructureConfig {
pub sample_rate: usize,
pub input_pad: usize,
pub stft: Conv1dConfig,
pub encoder: ConvSeq1dConfig,
pub lstm: FusedLstmConfig,
pub decoder: Conv1dConfig,
}
impl SileroVadMeta for SileroVadStructureConfig {
fn sample_rate(&self) -> usize {
self.sample_rate
}
fn n_freq(&self) -> usize {
self.stft.channels_out / 2
}
fn input_pad(&self) -> usize {
self.input_pad
}
fn stft_kernel(&self) -> usize {
self.stft.kernel_size
}
fn stft_stride(&self) -> usize {
self.stft.stride
}
fn d_hidden(&self) -> usize {
self.encoder.out_channels()
}
fn d_bottleneck(&self) -> usize {
self.encoder.blocks.last().unwrap().in_channels()
}
}
impl SileroVadStructureConfig {
pub fn validate(&self) -> BunsenResult<()> {
self.encoder.validate()?;
let hidden = self.d_hidden();
if self.encoder.in_channels() != self.n_freq() {
return Err(BunsenError::Invalid(format!(
"SileroVad encoder in_channels ({}) != n_freq ({})",
self.encoder.in_channels(),
self.n_freq(),
)));
}
if self.decoder.channels_in != hidden || self.decoder.channels_out != 1 {
return Err(BunsenError::Invalid(format!(
"SileroVad decoder must map hidden ({hidden}) -> 1, got {} -> {}",
self.decoder.channels_in, self.decoder.channels_out,
)));
}
Ok(())
}
}
impl<B: Backend> ModuleInit<B, SileroVad<B>> for SileroVadStructureConfig {
fn try_init(
&self,
device: &B::Device,
) -> BunsenResult<SileroVad<B>> {
self.validate()?;
Ok(SileroVad {
sample_rate: self.sample_rate,
input_pad: self.input_pad,
stft: self.stft.init(device),
encoder: self.encoder.try_init(device)?,
lstm: self.lstm.init(device),
decoder: self.decoder.init(device),
})
}
}
#[derive(Module, Debug)]
pub struct SileroVad<B: Backend> {
sample_rate: usize,
input_pad: usize,
pub stft: Conv1d<B>,
pub encoder: ConvSeq1d<B>,
pub lstm: FusedLstm<B>,
pub decoder: Conv1d<B>,
}
impl<B: Backend> SileroVadMeta for SileroVad<B> {
fn sample_rate(&self) -> usize {
self.sample_rate
}
fn n_freq(&self) -> usize {
self.stft.weight.dims()[0] / 2
}
fn input_pad(&self) -> usize {
self.input_pad
}
fn stft_kernel(&self) -> usize {
self.stft.kernel_size
}
fn stft_stride(&self) -> usize {
self.stft.stride
}
fn d_hidden(&self) -> usize {
self.encoder.out_channels()
}
fn d_bottleneck(&self) -> usize {
self.encoder.blocks.last().unwrap().in_channels()
}
}
impl<B: Backend> SileroVad<B> {
pub fn init_state(
&self,
batch: usize,
device: &B::Device,
) -> Tensor<B, 3> {
Tensor::zeros([2, batch, self.d_hidden()], device)
}
pub fn init_context(
&self,
batch: usize,
context_size: usize,
device: &B::Device,
) -> SileroVadContext<B> {
SileroVadContext {
sample_rate: self.sample_rate(),
context: Tensor::zeros([batch, context_size], device),
state: self.init_state(batch, device),
}
}
pub fn context_forward_sequence(
&self,
chunk_seq: Tensor<B, 3>,
context: SileroVadContext<B>,
) -> (Tensor<B, 2>, SileroVadContext<B>) {
let SileroVadContext {
sample_rate,
context,
state,
} = context;
assert_eq!(sample_rate, self.sample_rate());
cfg_select! {
any(test, debug_assertions) => {
use crate::contracts::{unpack_shape_contract, assert_shape_contract_periodically};
let [steps, batch] = unpack_shape_contract!(
["steps", "batch", "samples"],
&chunk_seq,
&["steps", "batch"],
&[("samples", self.chunk_size())]
);
let [context_size] = unpack_shape_contract!(
["batch", "context_size"],
&context,
&["context_size"],
&[("batch", batch)],
);
assert_shape_contract_periodically!(
[2, "batch", "d_hidden"],
&state,
&[("batch", batch), ("d_hidden", self.d_hidden())]
);
}
_ => {
let steps = chunk_seq.dims()[0];
let context_size = context.dims()[1];
}
}
let context: Tensor<B, 3> = context.unsqueeze_dim(0);
let context: Tensor<B, 3> = if steps <= 1 {
context
} else {
let tails = chunk_seq
.clone()
.slice(s![0..-1, .., -(context_size as isize)..]);
Tensor::cat(vec![context, tails], 0)
};
let ext_chunk_seq: Tensor<B, 3> = Tensor::cat(vec![context, chunk_seq.clone()], 2);
let context = ext_chunk_seq
.clone()
.slice(s![-1, .., -(context_size as isize)..])
.squeeze_dim::<2>(0);
let (out, state) = self.forward_sequence(ext_chunk_seq, state);
(
out,
SileroVadContext {
sample_rate,
context,
state,
},
)
}
pub fn context_forward(
&self,
chunk: Tensor<B, 2>,
context: SileroVadContext<B>,
) -> (Tensor<B, 1>, SileroVadContext<B>) {
let SileroVadContext {
sample_rate,
context,
state,
} = context;
assert_eq!(sample_rate, self.sample_rate());
cfg_select! {
any(test, debug_assertions) => {
use crate::contracts::{unpack_shape_contract, assert_shape_contract_periodically};
let [batch] = unpack_shape_contract!(
[ "batch", "samples"],
&chunk,
&["batch"],
&[("samples", self.chunk_size())]
);
let [context_size] = unpack_shape_contract!(
["batch", "context_size"],
&context,
&["context_size"],
&[("batch", batch)],
);
assert_shape_contract_periodically!(
[2, "batch", "d_hidden"],
&state,
&[("batch", batch), ("d_hidden", self.d_hidden())]
);
}
_ => {
let context_size = context.dims()[1];
}
}
let ext_input = Tensor::cat(vec![context, chunk], 1);
let context = ext_input.clone().slice(s![.., -(context_size as isize)..]);
let (out, state) = self.forward(ext_input, state);
(
out,
SileroVadContext {
sample_rate,
context,
state,
},
)
}
pub fn forward_sequence(
&self,
chunk_seq: Tensor<B, 3>,
state: Tensor<B, 3>,
) -> (Tensor<B, 2>, Tensor<B, 3>) {
cfg_select! {
any(test, debug_assertions) => {
let [steps, batch] = crate::contracts::unpack_shape_contract!(
["steps", "batch", "samples"],
&chunk_seq,
&["steps", "batch"],
);
crate::contracts::assert_shape_contract_periodically!(
[2, "batch", "d_hidden"],
&state,
&[("batch", batch), ("d_hidden", self.d_hidden())]
);
}
_ => {
let [steps, batch, _] = chunk_seq.dims();
}
}
let mut seq_features = self.frame_features(chunk_seq.flatten::<2>(0, 1)).reshape([
steps,
batch,
self.d_hidden(),
]);
let (mut hidden, mut cell) = Self::unpack_state(state);
macro_rules! process_steps {
(mut $acc:ident) => {{
for step in 0..steps {
let features = seq_features.clone().slice_dim(0, step).squeeze_dim::<2>(0);
(hidden, cell) = self.lstm_step(features, hidden, cell);
let step_hidden = hidden.clone().unsqueeze_dim::<3>(0);
$acc = $acc.slice_assign(s![step, .., ..], step_hidden);
}
$acc
}};
}
let seq_hidden: Tensor<B, 3> = if B::ad_enabled(&seq_features.device()) {
let mut seq_hidden = Tensor::zeros_like(&seq_features);
process_steps!(mut seq_hidden)
} else {
process_steps!(mut seq_features)
};
let out = self
.output_head(seq_hidden.flatten(0, 1))
.reshape([steps, batch]);
let state = Self::pack_state(hidden, cell);
(out, state)
}
pub fn forward(
&self,
chunk: Tensor<B, 2>,
state: Tensor<B, 3>,
) -> (Tensor<B, 1>, Tensor<B, 3>) {
#[cfg(any(test, debug_assertions))]
{
let [batch] =
crate::contracts::unpack_shape_contract!(["batch", "samples"], &chunk, &["batch"],);
crate::contracts::assert_shape_contract_periodically!(
[2, "batch", "d_hidden"],
&state,
&[("batch", batch), ("d_hidden", self.d_hidden())]
);
}
let features = self.frame_features(chunk);
let (hidden, cell) = Self::unpack_state(state);
let (hidden, cell) = self.lstm_step(features, hidden, cell);
(
self.output_head(hidden.clone()),
Self::pack_state(hidden, cell),
)
}
pub fn frame_features(
&self,
input: Tensor<B, 2>,
) -> Tensor<B, 2> {
#[cfg(any(test, debug_assertions))]
let [batch] =
crate::contracts::unpack_shape_contract!(["batch", "samples"], &input, &["batch"]);
let x: Tensor<B, 3> = input
.pad([(0, self.input_pad)], PadMode::Reflect)
.unsqueeze_dim::<3>(1);
let x = self.stft.forward(x);
#[cfg(any(test, debug_assertions))]
crate::contracts::assert_shape_contract_periodically!(
["batch", 2 * "n_freq", "T"],
&x,
&[("batch", batch), ("n_freq", self.n_freq())],
);
let [real_2, imag_2] = x.square().chunk(2, 1).try_into().unwrap();
let mag = (real_2 + imag_2).sqrt();
let x = self
.encoder
.forward(mag)
.slice_dim(2, 0)
.squeeze_dim::<2>(2);
#[cfg(any(test, debug_assertions))]
crate::contracts::assert_shape_contract_periodically!(
["batch", "d_hidden"],
&x,
&[("batch", batch), ("d_hidden", self.d_hidden())],
);
x
}
pub fn unpack_state(state: Tensor<B, 3>) -> (Tensor<B, 2>, Tensor<B, 2>) {
let [hidden, cell] = state.chunk(2, 0).try_into().unwrap();
(hidden.squeeze_dim::<2>(0), cell.squeeze_dim::<2>(0))
}
pub fn pack_state(
hidden: Tensor<B, 2>,
cell: Tensor<B, 2>,
) -> Tensor<B, 3> {
Tensor::stack(vec![hidden, cell], 0)
}
pub fn lstm_step(
&self,
features: Tensor<B, 2>,
hidden: Tensor<B, 2>,
cell: Tensor<B, 2>,
) -> (Tensor<B, 2>, Tensor<B, 2>) {
self.lstm.step(features, hidden, cell)
}
pub fn output_head(
&self,
hidden: Tensor<B, 2>,
) -> Tensor<B, 1> {
let x: Tensor<B, 3> = hidden.unsqueeze_dim::<3>(2);
let x = relu(x);
let x = self.decoder.forward(x);
let x = sigmoid(x);
let x = x.squeeze_dim::<2>(1);
let x = x.mean_dim(1);
x.squeeze_dim::<1>(1)
}
}
#[cfg(test)]
mod tests {
use burn::tensor::{
Distribution,
Tolerance,
backend::BackendTypes,
};
use super::*;
use crate::support::testing::PerformanceBackend;
type B = PerformanceBackend;
#[test]
fn test_config_meta() {
{
let cfg = SileroVadSignalConfig::standard_16khz();
assert_eq!(cfg.sample_rate, 16000);
assert_eq!(cfg.n_freq, 129);
let cfg = cfg.to_stft();
assert_eq!(cfg.sample_rate, 16000);
assert_eq!(cfg.n_freq, 129);
assert_eq!(cfg.stft_stride, 128);
assert_eq!(cfg.stft_kernel, 256);
assert_eq!(cfg.input_pad, 64);
assert_eq!(cfg.d_hidden, 128);
assert_eq!(cfg.d_bottleneck, 64);
let cfg = cfg.to_structure();
assert_eq!(cfg.sample_rate(), 16000);
assert_eq!(cfg.n_freq(), 129);
assert_eq!(cfg.chunk_size(), 512);
assert_eq!(cfg.input_pad(), 64);
assert_eq!(cfg.stft_kernel(), 256);
assert_eq!(cfg.stft_stride(), 128);
assert_eq!(cfg.gate_size(), 512);
assert_eq!(cfg.encoder.in_channels(), cfg.n_freq());
assert_eq!(cfg.d_hidden(), 128);
assert_eq!(cfg.d_bottleneck(), 64);
cfg.validate().unwrap();
}
{
let cfg = SileroVadSignalConfig::standard_8khz();
assert_eq!(cfg.sample_rate, 8000);
assert_eq!(cfg.n_freq, 65);
let cfg = cfg.to_stft();
assert_eq!(cfg.sample_rate, 8000);
assert_eq!(cfg.n_freq, 65);
assert_eq!(cfg.stft_stride, 64);
assert_eq!(cfg.stft_kernel, 128);
assert_eq!(cfg.input_pad, 32);
assert_eq!(cfg.d_hidden, 128);
assert_eq!(cfg.d_bottleneck, 64);
let cfg = cfg.to_structure();
assert_eq!(cfg.sample_rate(), 8000);
assert_eq!(cfg.n_freq(), 65);
assert_eq!(cfg.chunk_size(), 256);
assert_eq!(cfg.input_pad(), 32);
assert_eq!(cfg.stft_kernel(), 128);
assert_eq!(cfg.stft_stride(), 64);
assert_eq!(cfg.gate_size(), 512);
assert_eq!(cfg.encoder.in_channels(), cfg.n_freq());
assert_eq!(cfg.d_hidden(), 128);
assert_eq!(cfg.d_bottleneck(), 64);
cfg.validate().unwrap();
}
}
#[test]
fn test_validate_rejects_mismatch() {
let bad = SileroVadStructureConfig {
encoder: encoder_config(64, 128, 64),
..SileroVadSignalConfig::standard_16khz().to_structure()
};
assert!(matches!(bad.validate(), Err(BunsenError::Invalid(_))));
}
#[test]
fn test_config_meta_matches_module() {
let device = Default::default();
for (cfg, n_freq, chunk_size) in [
(
SileroVadSignalConfig::standard_16khz().to_structure(),
129,
512,
),
(
SileroVadSignalConfig::standard_8khz().to_structure(),
65,
256,
),
] {
assert_eq!(cfg.chunk_size(), chunk_size);
assert_eq!(cfg.n_freq(), n_freq);
let model: SileroVad<B> = cfg.init(&device);
assert_eq!(model.sample_rate(), cfg.sample_rate());
assert_eq!(model.chunk_size(), cfg.chunk_size());
assert_eq!(model.d_hidden(), cfg.d_hidden());
assert_eq!(model.input_pad(), cfg.input_pad());
assert_eq!(model.stft_kernel(), cfg.stft_kernel());
assert_eq!(model.stft_stride(), cfg.stft_stride());
assert_eq!(model.gate_size(), cfg.gate_size());
assert_eq!(model.d_bottleneck(), cfg.d_bottleneck());
}
}
#[test]
#[serial_test::serial]
fn test_forward_shapes_and_range() {
let device = Default::default();
for cfg in [
SileroVadSignalConfig::standard_16khz().to_structure(),
SileroVadSignalConfig::standard_8khz().to_structure(),
] {
let model: SileroVad<B> = cfg.init(&device);
let batch = 3;
let context = 64;
let input = Tensor::<B, 2>::random(
[batch, context + model.chunk_size()],
Distribution::Default,
&device,
);
let state = model.init_state(batch, &device);
let (prob, next_state) = model.forward(input, state);
assert_eq!(prob.dims(), [batch]);
assert_eq!(next_state.dims(), [2, batch, 128]);
let probs: Vec<f32> = prob.into_data().to_vec().unwrap();
assert!(probs.iter().all(|&p| (0.0..=1.0).contains(&p)));
}
}
#[test]
#[serial_test::serial]
fn test_forward_sequence_shapes() {
let device = Default::default();
let batch = 8;
let steps = 5;
let context = 64;
for cfg in [
SileroVadSignalConfig::standard_16khz().to_structure(),
SileroVadSignalConfig::standard_8khz().to_structure(),
] {
let model: SileroVad<B> = cfg.init(&device);
let input = Tensor::random(
[steps, batch, context + model.chunk_size()],
Distribution::Default,
&device,
);
let state = model.init_state(batch, &device);
let (probs, next_state) = model.forward_sequence(input, state);
assert_eq!(probs.dims(), [steps, batch]);
assert_eq!(next_state.dims(), [2, batch, 128]);
}
}
fn check_sequence_matches_stepwise<B: Backend, F>()
where
F: num_traits::Float + burn::tensor::Element,
{
let device = Default::default();
let model: SileroVad<B> = SileroVadSignalConfig::standard_16khz()
.to_structure()
.init(&device);
let steps = 5;
let batch = 8;
let context = 64;
let input = Tensor::random(
[steps, batch, context + model.chunk_size()],
Distribution::Default,
&device,
);
let mut state = model.init_state(batch, &device);
let (seq_probs, seq_state) = model.forward_sequence(input.clone(), state.clone());
let mut step_probs = Vec::with_capacity(steps);
for step in 0..steps {
let chunk = input.clone().slice_dim(0, step).squeeze_dim::<2>(0);
let (prob, next_state) = model.forward(chunk, state);
state = next_state;
step_probs.push(prob);
}
let step_probs: Tensor<B, 2> = Tensor::stack(step_probs, 0);
let tol = Tolerance::<F>::default();
seq_probs
.into_data()
.assert_approx_eq::<F>(&step_probs.into_data(), tol);
seq_state
.into_data()
.assert_approx_eq::<F>(&state.into_data(), tol);
}
#[test]
#[serial_test::serial]
fn test_sequence_matches_stepwise_no_ad() {
type F = <B as BackendTypes>::FloatElem;
check_sequence_matches_stepwise::<B, F>();
}
#[test]
#[serial_test::serial]
fn test_sequence_matches_stepwise_autodiff() {
type F = <B as BackendTypes>::FloatElem;
check_sequence_matches_stepwise::<burn::backend::Autodiff<B>, F>();
}
}