use burn::{
Tensor,
config::Config,
module::Module,
prelude::Backend,
};
use crate::kits::speech::silero_vad::SileroVad;
pub trait SileroVadContextMeta {
fn sample_rate(&self) -> usize;
fn batch_size(&self) -> usize;
fn context_size(&self) -> usize;
}
#[derive(Config, Debug)]
pub struct SileroVadContextConfig {
pub sample_rate: usize,
#[config(default = "1")]
pub batch_size: usize,
#[config(default = "64")]
pub context_size: usize,
}
impl SileroVadContextMeta for SileroVadContextConfig {
fn sample_rate(&self) -> usize {
self.sample_rate
}
fn batch_size(&self) -> usize {
self.batch_size
}
fn context_size(&self) -> usize {
self.context_size
}
}
#[derive(Module, Debug)]
pub struct SileroVadContext<B: Backend> {
pub sample_rate: usize,
pub context: Tensor<B, 2>,
pub state: Tensor<B, 3>,
}
impl<B: Backend> SileroVadContextMeta for SileroVadContext<B> {
fn sample_rate(&self) -> usize {
self.sample_rate
}
fn batch_size(&self) -> usize {
self.context.dims()[0]
}
fn context_size(&self) -> usize {
self.context.dims()[1]
}
}
impl SileroVadContextConfig {
pub fn init<B: Backend>(
&self,
vad: &SileroVad<B>,
device: &B::Device,
) -> SileroVadContext<B> {
vad.init_context(self.batch_size, self.context_size, device)
}
}