bunsen 0.25.0

bunsen is a batteries included common library for burn
Documentation
//! Pretrained models

use std::path::{
    Path,
    PathBuf,
};

use burn::prelude::Backend;
use burn_store::{
    BurnpackStore,
    ModuleSnapshot,
};

use crate::{
    burner::module::ModuleInit,
    errors::{
        BunsenError,
        BunsenResult,
    },
    kits::speech::silero_vad::{
        SileroVad,
        SileroVadConfig,
    },
};

/// Load a pretrained Silero VAD model from a file.
pub fn load_pretrained_silerovad<B: Backend, P: AsRef<Path>>(
    path: P,
    device: &B::Device,
) -> BunsenResult<SileroVad<B>> {
    let path = path.as_ref();
    let path: PathBuf = path.to_path_buf();

    let vad16: SileroVad<B> = {
        let cfg = SileroVadConfig::standard_16khz();

        let mut store = BurnpackStore::from_file(path.clone())
            .with_remap_pattern("conv1d37", "stft")
            .with_remap_pattern("conv1d38", "encoder.blocks.0.conv")
            .with_remap_pattern("conv1d39", "encoder.blocks.1.conv")
            .with_remap_pattern("conv1d40", "encoder.blocks.2.conv")
            .with_remap_pattern("conv1d41", "encoder.blocks.3.conv")
            .with_remap_pattern("linear13", "input_gate")
            .with_remap_pattern("linear14", "hidden_gate")
            .with_remap_pattern("conv1d42", "decoder");

        // println!("keys: {:#?}", store.keys());

        let mut module = cfg.try_init(device)?;
        module
            .load_from(&mut store)
            .map_err(BunsenError::external)?;

        module
    };

    let _vad8: SileroVad<B> = {
        let cfg = SileroVadConfig::standard_8khz();

        let mut store = BurnpackStore::from_file(path.clone())
            .with_remap_pattern("conv1d43", "stft")
            .with_remap_pattern("conv1d44", "encoder.blocks.0.conv")
            .with_remap_pattern("conv1d45", "encoder.blocks.1.conv")
            .with_remap_pattern("conv1d46", "encoder.blocks.2.conv")
            .with_remap_pattern("conv1d47", "encoder.blocks.3.conv")
            .with_remap_pattern("linear15", "input_gate")
            .with_remap_pattern("linear16", "hidden_gate")
            .with_remap_pattern("conv1d48", "decoder");

        // println!("keys: {:#?}", store.keys());

        let mut module = cfg.try_init(device)?;
        module
            .load_from(&mut store)
            .map_err(BunsenError::external)?;

        module
    };

    Ok(vad16)
}

#[cfg(test)]
mod tests {
    use std::path::PathBuf;

    use super::*;
    use crate::{
        errors::*,
        support::testing::CpuBackend,
    };

    #[test]
    #[ignore]
    fn test_load_pretrained() {
        type B = CpuBackend;

        let path = PathBuf::from(
            "/home/crutcher/git/fast-whisper-burn/src/vad/silero_vad_op18_ifless.bpk",
        );
        let device = Default::default();

        let _vad: SileroVad<B> = load_pretrained_silerovad(path, &device).ok_or_panic();
    }
}