#![allow(clippy::too_many_arguments)]
pub mod data;
#[cfg(all(feature = "ml_sample_gather", feature = "analyze_mic"))]
pub mod gather;
pub mod helpers;
pub mod model;
pub mod precision;
#[cfg(feature = "ml_sample_process")]
pub mod process;
pub use precision::StorePrecisionSettings;
#[cfg(feature = "ml_train")]
pub use precision::{PrecisionElement, PRECISION_DTYPE};
use burn::config::Config;
use std::path::PathBuf;
pub const FREQUENCY_SPACE_SIZE: usize = 8192;
pub const NOTE_SIGNATURE_SIZE: usize = 128;
pub const PITCH_CLASS_COUNT: usize = 12;
pub const DETERMINISTIC_GUESS_SIZE: usize = NOTE_SIGNATURE_SIZE;
pub const MEL_SPACE_SIZE: usize = 512;
#[cfg(any(
all(
feature = "ml_loader_note_binned_convolution",
any(feature = "ml_loader_mel", feature = "ml_loader_frequency", feature = "ml_loader_frequency_pooled")
),
all(feature = "ml_loader_mel", any(feature = "ml_loader_frequency", feature = "ml_loader_frequency_pooled")),
all(feature = "ml_loader_frequency", feature = "ml_loader_frequency_pooled"),
))]
compile_error!(
"Multiple ml_loader_* features enabled; enable exactly one of: \
ml_loader_note_binned_convolution, ml_loader_mel, ml_loader_frequency, ml_loader_frequency_pooled."
);
#[cfg(not(any(
feature = "ml_loader_note_binned_convolution",
feature = "ml_loader_mel",
feature = "ml_loader_frequency",
feature = "ml_loader_frequency_pooled",
)))]
compile_error!(
"No ml_loader_* feature enabled; enable exactly one of: \
ml_loader_note_binned_convolution, ml_loader_mel, ml_loader_frequency, ml_loader_frequency_pooled."
);
#[cfg(feature = "ml_loader_note_binned_convolution")]
const INPUT_BASE_SIZE: usize = NOTE_SIGNATURE_SIZE;
#[cfg(feature = "ml_loader_mel")]
const INPUT_BASE_SIZE: usize = MEL_SPACE_SIZE;
#[cfg(feature = "ml_loader_frequency")]
const INPUT_BASE_SIZE: usize = FREQUENCY_SPACE_SIZE;
#[cfg(feature = "ml_loader_frequency_pooled")]
pub const FREQUENCY_POOL_FACTOR: usize = 16;
#[cfg(feature = "ml_loader_frequency_pooled")]
pub const FREQUENCY_SPACE_POOLED_SIZE: usize = FREQUENCY_SPACE_SIZE / FREQUENCY_POOL_FACTOR;
#[cfg(feature = "ml_loader_frequency_pooled")]
const INPUT_BASE_SIZE: usize = FREQUENCY_SPACE_POOLED_SIZE;
#[cfg(feature = "ml_loader_include_deterministic_guess")]
pub const INPUT_SPACE_SIZE: usize = INPUT_BASE_SIZE + DETERMINISTIC_GUESS_SIZE;
#[cfg(not(feature = "ml_loader_include_deterministic_guess"))]
pub const INPUT_SPACE_SIZE: usize = INPUT_BASE_SIZE;
#[cfg(not(any(feature = "ml_target_full", feature = "ml_target_folded", feature = "ml_target_folded_bass")))]
compile_error!("No ml_target_* feature enabled; enable exactly one of: ml_target_full, ml_target_folded, ml_target_folded_bass.");
#[cfg(any(
all(feature = "ml_target_full", feature = "ml_target_folded"),
all(feature = "ml_target_full", feature = "ml_target_folded_bass"),
all(feature = "ml_target_folded", feature = "ml_target_folded_bass"),
))]
compile_error!("Multiple ml_target_* features enabled; select exactly one of: ml_target_full, ml_target_folded, ml_target_folded_bass.");
#[cfg(feature = "ml_target_full")]
pub const TARGET_SPACE_SIZE: usize = NOTE_SIGNATURE_SIZE;
#[cfg(feature = "ml_target_folded")]
pub const TARGET_SPACE_SIZE: usize = PITCH_CLASS_COUNT;
#[cfg(feature = "ml_target_folded_bass")]
pub const TARGET_SPACE_SIZE: usize = 2 * PITCH_CLASS_COUNT;
pub const NUM_CLASSES: usize = TARGET_SPACE_SIZE;
#[cfg(feature = "ml_target_folded_bass")]
pub const TARGET_FOLDED_BASS_OFFSET: usize = 0;
#[cfg(feature = "ml_target_folded_bass")]
pub const TARGET_FOLDED_BASS_NOTE_OFFSET: usize = PITCH_CLASS_COUNT;
#[derive(Debug, Config)]
pub struct TrainConfig {
pub noise_asset_root: String,
pub training_sources: Vec<String>,
pub validation_sources: Vec<String>,
pub destination: String,
pub log: String,
pub simulation_size: usize,
pub simulation_peak_radius: f32,
pub simulation_harmonic_decay: f32,
pub simulation_frequency_wobble: f32,
pub captured_oversample_factor: usize,
pub mha_heads: usize,
pub dropout: f64,
pub trunk_hidden_size: usize,
pub model_epochs: usize,
pub model_batch_size: usize,
pub model_workers: usize,
pub model_seed: u64,
pub adam_learning_rate: f64,
pub adam_weight_decay: f32,
pub adam_beta1: f32,
pub adam_beta2: f32,
pub adam_epsilon: f32,
pub no_plots: bool,
}
#[derive(Clone, Debug)]
pub struct KordItem {
pub path: PathBuf,
pub frequency_space: [f32; FREQUENCY_SPACE_SIZE],
pub label: u128,
}
impl Default for KordItem {
fn default() -> Self {
Self {
path: PathBuf::new(),
frequency_space: [0.0; FREQUENCY_SPACE_SIZE],
label: 0,
}
}
}