qualia_core_db/wgsl_forge/audio/
mel.rs1use crate::wgsl_forge::ForgeError;
22
23pub const MEL_APPLY_WGSL: &str = include_str!("../../shaders/audio_mel.wgsl");
25pub const MEL_APPLY_ENTRY: &str = "mel_apply";
27
28pub fn mel_apply_cpu(
33 spectrum: &[f32],
34 mel_fb: &[f32],
35 n_frames: usize,
36 n_bins: usize,
37 n_mel: usize,
38) -> Vec<f32> {
39 let mut out = vec![0.0f32; n_frames * n_mel];
40 for frame in 0..n_frames {
41 let spec_base = frame * n_bins;
42 for m in 0..n_mel {
43 let fb_base = m * n_bins;
44 let mut acc = 0.0f32;
45 for b in 0..n_bins {
46 acc += spectrum[spec_base + b] * mel_fb[fb_base + b];
47 }
48 out[frame * n_mel + m] = acc;
49 }
50 }
51 out
52}
53
54pub fn mel_apply_forge(
69 spectrum: &[f32],
70 mel_fb: &[f32],
71 n_frames: usize,
72 n_bins: usize,
73 n_mel: usize,
74) -> Result<Vec<f32>, ForgeError> {
75 use crate::wgsl_forge::execute::{
76 BindingUsage, QualiaCompute, WgpuComputeContext, WgpuPipeline,
77 };
78 use crate::wgsl_forge::Schedule;
79
80 if n_frames == 0 || n_bins == 0 || n_mel == 0 {
81 return Err(ForgeError::GpuValidation(format!(
82 "mel_apply_forge: dimensions must be non-zero (n_frames={n_frames}, n_bins={n_bins}, n_mel={n_mel})"
83 )));
84 }
85 if spectrum.len() != n_frames * n_bins {
86 return Err(ForgeError::GpuValidation(format!(
87 "mel_apply_forge: spectrum length {} != n_frames*n_bins {}",
88 spectrum.len(),
89 n_frames * n_bins
90 )));
91 }
92 if mel_fb.len() != n_mel * n_bins {
93 return Err(ForgeError::GpuValidation(format!(
94 "mel_apply_forge: mel_fb length {} != n_mel*n_bins {}",
95 mel_fb.len(),
96 n_mel * n_bins
97 )));
98 }
99
100 let out_len = n_frames * n_mel;
101 let total_floats = spectrum.len() + mel_fb.len() + out_len + 4;
102 let capacity = (total_floats * 4).max(4 << 20);
103 let shared = crate::gpu_context::device_registry::try_auxiliary_gpu().ok_or_else(|| {
107 ForgeError::GpuUnavailable(
108 "mel_apply_forge: no GPU circuit available (auxiliary→primary both absent)".to_string(),
109 )
110 })?;
111 let mut ctx = WgpuComputeContext::from_device(
112 shared.device.clone(),
113 shared.queue.clone(),
114 &shared.adapter_caps,
115 capacity,
116 )?;
117
118 let view_spectrum = ctx.allocate_and_write(
119 bytemuck::cast_slice(spectrum),
120 0,
121 0,
122 BindingUsage::StorageRead,
123 )?;
124 let view_fb = ctx.allocate_and_write(
125 bytemuck::cast_slice(mel_fb),
126 1,
127 0,
128 BindingUsage::StorageRead,
129 )?;
130 let zeros = vec![0.0f32; out_len];
131 let view_out = ctx.allocate_and_write(
132 bytemuck::cast_slice(&zeros),
133 2,
134 0,
135 BindingUsage::StorageReadWrite,
136 )?;
137 let params = [n_frames as f32, n_bins as f32, n_mel as f32];
138 let view_params = ctx.allocate_and_write(
139 bytemuck::cast_slice(¶ms),
140 3,
141 0,
142 BindingUsage::StorageRead,
143 )?;
144
145 let buffers = vec![view_spectrum, view_fb, view_out, view_params];
146 let pipeline = WgpuPipeline::compile(&ctx, MEL_APPLY_WGSL, MEL_APPLY_ENTRY)?;
147 let schedule = Schedule {
148 workgroup_size: 64,
149 ..Default::default()
150 };
151 pipeline.dispatch(&buffers, &schedule, out_len)?;
152 let mut out = ctx.read_buffer_f32(&view_out)?;
153 out.truncate(out_len);
154 Ok(out)
155}
156
157pub fn mel_apply(
164 spectrum: &[f32],
165 mel_fb: &[f32],
166 n_frames: usize,
167 n_bins: usize,
168 n_mel: usize,
169) -> Vec<f32> {
170 if crate::wgsl_forge::dispatch::caps().wgpu {
171 if let Ok(out) = mel_apply_forge(spectrum, mel_fb, n_frames, n_bins, n_mel) {
172 return out;
173 }
174 }
177 mel_apply_cpu(spectrum, mel_fb, n_frames, n_bins, n_mel)
178}
179
180#[cfg(test)]
181mod tests {
182 use super::*;
183 use crate::wgsl_forge::validate::validate_wgsl;
184
185 #[test]
188 fn mel_apply_wgsl_validates() {
189 let report = validate_wgsl(MEL_APPLY_WGSL).expect("mel-apply WGSL must naga-validate");
190 assert!(
191 report.entry_points.iter().any(|e| e == MEL_APPLY_ENTRY),
192 "validated module must expose {MEL_APPLY_ENTRY}; got {:?}",
193 report.entry_points
194 );
195 }
196
197 #[test]
200 fn mel_apply_cpu_matches_reference() {
201 let spectrum = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
205 let mel_fb = vec![1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 1.0];
209
210 let out = mel_apply_cpu(&spectrum, &mel_fb, 2, 4, 2);
211
212 assert_eq!(out, vec![3.0, 7.0, 11.0, 15.0]);
215 }
216
217 #[test]
220 fn mel_apply_public_matches_cpu() {
221 let (n_frames, n_bins, n_mel) = (3usize, 5usize, 4usize);
222 let spectrum: Vec<f32> = (0..n_frames * n_bins)
223 .map(|k| (k as f32) * 0.5 - 3.0)
224 .collect();
225 let mut mel_fb = vec![0.0f32; n_mel * n_bins];
227 for m in 0..n_mel {
228 for b in 0..n_bins {
229 mel_fb[m * n_bins + b] = ((m + b) as f32 % 3.0) * 0.25;
230 }
231 }
232 let public = mel_apply(&spectrum, &mel_fb, n_frames, n_bins, n_mel);
233 let oracle = mel_apply_cpu(&spectrum, &mel_fb, n_frames, n_bins, n_mel);
234 assert_eq!(public.len(), oracle.len());
235 for (p, o) in public.iter().zip(oracle.iter()) {
236 let tol = 1e-3 * o.abs().max(1.0);
237 assert!((p - o).abs() <= tol, "public/CPU mismatch: {p} vs {o}");
238 }
239 }
240
241 #[test]
244 #[serial_test::serial(gpu)]
245 fn mel_gpu_matches_oracle() {
246 if !crate::wgsl_forge::test_gpu_available() {
247 return;
248 }
249 if let Some(shared) = crate::gpu_context::device_registry::try_auxiliary_gpu() {
252 let caps = &shared.adapter_caps;
253 eprintln!(
254 "mel_gpu_matches_oracle: forge on adapter '{}' ({:?}, {:?})",
255 caps.name, caps.device_type, caps.backend
256 );
257 }
258 let (n_frames, n_bins, n_mel) = (17usize, 33usize, 12usize);
259 let spectrum: Vec<f32> = (0..n_frames * n_bins)
260 .map(|k| ((k as f32) * 0.017).sin().abs() + 0.001)
261 .collect();
262 let mut mel_fb = vec![0.0f32; n_mel * n_bins];
264 for m in 0..n_mel {
265 let centre = (m as f32 + 1.0) * (n_bins as f32) / (n_mel as f32 + 1.0);
266 for b in 0..n_bins {
267 let w = 1.0 - ((b as f32 - centre).abs() / 3.0);
268 mel_fb[m * n_bins + b] = w.max(0.0);
269 }
270 }
271 let expected = mel_apply_cpu(&spectrum, &mel_fb, n_frames, n_bins, n_mel);
272 let gpu = mel_apply_forge(&spectrum, &mel_fb, n_frames, n_bins, n_mel)
273 .expect("mel_apply_forge on an available device");
274 assert_eq!(gpu.len(), expected.len());
275 for (g, e) in gpu.iter().zip(expected.iter()) {
276 let tol = 1e-3 * e.abs().max(1.0);
277 assert!((g - e).abs() <= tol, "GPU/CPU mismatch: {g} vs {e}");
278 }
279 }
280}