1#![allow(clippy::too_many_arguments)]
5
6pub mod data;
7#[cfg(all(feature = "ml_sample_gather", feature = "analyze_mic"))]
8pub mod gather;
9pub mod helpers;
10pub mod model;
11pub mod precision;
12#[cfg(feature = "ml_sample_process")]
13pub mod process;
14
15pub use precision::StorePrecisionSettings;
16#[cfg(feature = "ml_train")]
17pub use precision::{PrecisionElement, PRECISION_DTYPE};
18
19use burn::config::Config;
20use std::path::PathBuf;
21
22pub const FREQUENCY_SPACE_SIZE: usize = 8192;
26
27pub const NOTE_SIGNATURE_SIZE: usize = 128;
29
30pub const PITCH_CLASS_COUNT: usize = 12;
32
33pub const DETERMINISTIC_GUESS_SIZE: usize = NOTE_SIGNATURE_SIZE;
35
36pub const MEL_SPACE_SIZE: usize = 512;
38
39#[cfg(any(
41 all(
42 feature = "ml_loader_note_binned_convolution",
43 any(feature = "ml_loader_mel", feature = "ml_loader_frequency", feature = "ml_loader_frequency_pooled")
44 ),
45 all(feature = "ml_loader_mel", any(feature = "ml_loader_frequency", feature = "ml_loader_frequency_pooled")),
46 all(feature = "ml_loader_frequency", feature = "ml_loader_frequency_pooled"),
47))]
48compile_error!(
49 "Multiple ml_loader_* features enabled; enable exactly one of: \
50 ml_loader_note_binned_convolution, ml_loader_mel, ml_loader_frequency, ml_loader_frequency_pooled."
51);
52
53#[cfg(not(any(
54 feature = "ml_loader_note_binned_convolution",
55 feature = "ml_loader_mel",
56 feature = "ml_loader_frequency",
57 feature = "ml_loader_frequency_pooled",
58)))]
59compile_error!(
60 "No ml_loader_* feature enabled; enable exactly one of: \
61 ml_loader_note_binned_convolution, ml_loader_mel, ml_loader_frequency, ml_loader_frequency_pooled."
62);
63
64#[cfg(feature = "ml_loader_note_binned_convolution")]
66const INPUT_BASE_SIZE: usize = NOTE_SIGNATURE_SIZE;
67
68#[cfg(feature = "ml_loader_mel")]
70const INPUT_BASE_SIZE: usize = MEL_SPACE_SIZE;
71
72#[cfg(feature = "ml_loader_frequency")]
74const INPUT_BASE_SIZE: usize = FREQUENCY_SPACE_SIZE;
75
76#[cfg(feature = "ml_loader_frequency_pooled")]
78pub const FREQUENCY_POOL_FACTOR: usize = 16;
79
80#[cfg(feature = "ml_loader_frequency_pooled")]
82pub const FREQUENCY_SPACE_POOLED_SIZE: usize = FREQUENCY_SPACE_SIZE / FREQUENCY_POOL_FACTOR;
83
84#[cfg(feature = "ml_loader_frequency_pooled")]
86const INPUT_BASE_SIZE: usize = FREQUENCY_SPACE_POOLED_SIZE;
87
88#[cfg(feature = "ml_loader_include_deterministic_guess")]
90pub const INPUT_SPACE_SIZE: usize = INPUT_BASE_SIZE + DETERMINISTIC_GUESS_SIZE;
91
92#[cfg(not(feature = "ml_loader_include_deterministic_guess"))]
94pub const INPUT_SPACE_SIZE: usize = INPUT_BASE_SIZE;
95
96#[cfg(not(any(feature = "ml_target_full", feature = "ml_target_folded", feature = "ml_target_folded_bass")))]
98compile_error!("No ml_target_* feature enabled; enable exactly one of: ml_target_full, ml_target_folded, ml_target_folded_bass.");
99
100#[cfg(any(
101 all(feature = "ml_target_full", feature = "ml_target_folded"),
102 all(feature = "ml_target_full", feature = "ml_target_folded_bass"),
103 all(feature = "ml_target_folded", feature = "ml_target_folded_bass"),
104))]
105compile_error!("Multiple ml_target_* features enabled; select exactly one of: ml_target_full, ml_target_folded, ml_target_folded_bass.");
106
107#[cfg(feature = "ml_target_full")]
109pub const TARGET_SPACE_SIZE: usize = NOTE_SIGNATURE_SIZE;
110
111#[cfg(feature = "ml_target_folded")]
113pub const TARGET_SPACE_SIZE: usize = PITCH_CLASS_COUNT;
114
115#[cfg(feature = "ml_target_folded_bass")]
117pub const TARGET_SPACE_SIZE: usize = 2 * PITCH_CLASS_COUNT;
118
119pub const NUM_CLASSES: usize = TARGET_SPACE_SIZE;
121
122#[cfg(feature = "ml_target_folded_bass")]
124pub const TARGET_FOLDED_BASS_OFFSET: usize = 0;
125
126#[cfg(feature = "ml_target_folded_bass")]
128pub const TARGET_FOLDED_BASS_NOTE_OFFSET: usize = PITCH_CLASS_COUNT;
129
130#[derive(Debug, Config)]
134pub struct TrainConfig {
135 pub noise_asset_root: String,
137 pub training_sources: Vec<String>,
139 pub validation_sources: Vec<String>,
141 pub destination: String,
143 pub log: String,
145
146 pub simulation_size: usize,
148 pub simulation_peak_radius: f32,
150 pub simulation_harmonic_decay: f32,
152 pub simulation_frequency_wobble: f32,
154
155 pub captured_oversample_factor: usize,
157
158 pub mha_heads: usize,
160 pub dropout: f64,
162
163 pub trunk_hidden_size: usize,
165
166 pub model_epochs: usize,
168 pub model_batch_size: usize,
170 pub model_workers: usize,
172 pub model_seed: u64,
174
175 pub adam_learning_rate: f64,
177 pub adam_weight_decay: f32,
179 pub adam_beta1: f32,
181 pub adam_beta2: f32,
183 pub adam_epsilon: f32,
185
186 pub no_plots: bool,
188}
189
190#[derive(Clone, Debug)]
194pub struct KordItem {
195 pub path: PathBuf,
197 pub frequency_space: [f32; FREQUENCY_SPACE_SIZE],
199 pub label: u128,
201}
202
203impl Default for KordItem {
204 fn default() -> Self {
205 Self {
206 path: PathBuf::new(),
207 frequency_space: [0.0; FREQUENCY_SPACE_SIZE],
208 label: 0,
209 }
210 }
211}