/
/
1"""Smart Fades audio analysis provider."""
2
3from __future__ import annotations
4
5import asyncio
6import time
7from dataclasses import dataclass, field
8from datetime import timedelta
9from typing import TYPE_CHECKING, Any
10
11import numpy as np
12import soxr
13import torch
14from beat_this.inference import Spect2Frames, aggregate_prediction, split_piece
15from music_assistant_models.config_entries import ConfigEntry
16from music_assistant_models.enums import ConfigEntryType, MediaType
17from torchaudio.transforms import SpectralCentroid
18
19from music_assistant.constants import VERBOSE_LOG_LEVEL
20from music_assistant.helpers.datetime import utc
21from music_assistant.helpers.util import is_arm, system_meets_requirements
22from music_assistant.models.audio_analysis import AudioAnalysisData, AudioAnalysisError
23from music_assistant.models.audio_analysis_provider import (
24 ACCUMULATING_ANALYSIS_MAX_DURATION_SECONDS,
25 AudioAnalysisProvider,
26)
27
28from .dbn_postprocessor import DBNDownBeatTracker
29from .feature_extractor import AdvancedBeatFeatureExtractor
30from .helpers import (
31 aggregate_series_to_bins,
32 calculate_overall_bpm,
33 compute_band_rms_frames,
34 decode_pcm_chunk_to_mono,
35)
36from .resources.skey_model import KEY_MAP as SKEY_KEY_MAP
37from .resources.skey_model import VQT, ChromaNet, CropCQT, load_skey_components
38from .vocal_activity import (
39 FIRERED_SAMPLE_RATE,
40 FireRedFbank,
41 infer_firered_chunk,
42 load_firered_components,
43 split_firered_features,
44 vocal_activity_probabilities,
45)
46
47if TYPE_CHECKING:
48 import numpy.typing as npt
49 from music_assistant_models.config_entries import ProviderConfig
50 from music_assistant_models.enums import ProviderFeature
51 from music_assistant_models.media_items import AudioFormat
52 from music_assistant_models.provider import ProviderManifest
53 from music_assistant_models.streamdetails import StreamDetails
54
55 from music_assistant.mass import MusicAssistant
56
57ANALYSIS_SAMPLE_RATE = 22050
58# Below the recommended thresholds the provider still runs, but we surface an
59# informational notice (see get_config_entries) as it may be tight under load.
60RECOMMENDED_RAM_GB = 6.0
61RECOMMENDED_CPU_CORES = 4
62# Beat This predicts a long track as fixed windows. These are the values the model was trained
63# and released with (30s at 50 fps, plus the loss-border frames its predictions are unreliable
64# on), so a windowed prediction is identical to a whole-track one. Do not tune them: a window
65# of another length puts the model off its training distribution across the whole window.
66BEAT_WINDOW_FRAMES = 1500
67BEAT_WINDOW_BORDER_FRAMES = 6
68BEAT_WINDOW_OVERLAP_MODE = "keep_first"
69# While a player streams, wait this many times a window's own compute time before starting the
70# next one, so beat inference does not occupy a core continuously.
71BEAT_WINDOW_PACE_RATIO = 1.0
72# A model failure is often transient: an idle unload frees the models while an analysis that
73# outlived its own session is still running, or a one-off torch/hardware error. Record those
74# with this retry horizon instead of a permanent row that blocks the track forever.
75MODEL_FAILURE_RETRY_DELAY = timedelta(hours=24)
76
77
78@dataclass
79class SmartFadesData:
80 """Per-session data for smart fades analysis."""
81
82 item_id: str
83 provider: str
84 input_audio_format: AudioFormat
85 block_samples: int
86 features: AdvancedBeatFeatureExtractor
87 resampler: soxr.ResampleStream | None = None
88 pcm_buffer: list[np.ndarray] = field(default_factory=list)
89 pcm_samples: int = 0
90 total_pcm_samples: int = 0
91 beats_feature_blocks: list[np.ndarray] = field(default_factory=list)
92 energy_chunks: list[np.ndarray] = field(default_factory=list)
93 centroid_chunks: list[np.ndarray] = field(default_factory=list)
94 frequency_band_chunks: dict[str, list[np.ndarray]] = field(default_factory=dict)
95 musical_key_feature_blocks: list[torch.Tensor] = field(default_factory=list)
96 vocal_resampler: soxr.ResampleStream | None = None
97 vocal_fbank: FireRedFbank | None = None
98 vocal_feature_blocks: list[np.ndarray] = field(default_factory=list)
99
100
101class SmartFadesProvider(AudioAnalysisProvider):
102 """Smart fades audio analysis provider using Beat This for beat tracking."""
103
104 max_analysis_duration = ACCUMULATING_ANALYSIS_MAX_DURATION_SECONDS
105 # v3: FireRed AED vocal activity
106 analysis_version = 3
107 has_unloadable_models = True
108
109 def __init__(
110 self,
111 mass: MusicAssistant,
112 manifest: ProviderManifest,
113 config: ProviderConfig,
114 supported_features: set[ProviderFeature],
115 ) -> None:
116 """Initialize the provider."""
117 super().__init__(mass, manifest, config, supported_features)
118 self._data: dict[str, SmartFadesData] = {}
119 self._device = "cpu"
120 # Populated by _load_models and cleared again by _free_models, so every use goes
121 # through _require_loaded rather than assuming the models are resident.
122 self._beat_this_model: Spect2Frames | None = None
123 self._beat_this_post_processor: DBNDownBeatTracker | None = None
124 self._skey_vqt: VQT | None = None
125 self._skey_chromanet: ChromaNet | None = None
126 self._skey_crop: CropCQT | None = None
127 self._spectral_centroid: SpectralCentroid | None = None
128 self._firered_model: torch.nn.Module | None = None
129 self._firered_cmvn_means: npt.NDArray[np.float64] | None = None
130 self._firered_cmvn_inverse_std: npt.NDArray[np.float64] | None = None
131
132 async def get_config_entries(self) -> tuple[ConfigEntry, ...]:
133 """Return config entries for this provider."""
134 return (
135 ConfigEntry(
136 key="resource_warning",
137 type=ConfigEntryType.ALERT,
138 required=False,
139 hidden=system_meets_requirements(
140 min_memory_gb=RECOMMENDED_RAM_GB,
141 min_cpu_cores=RECOMMENDED_CPU_CORES,
142 ),
143 ),
144 )
145
146 async def handle_async_init(self) -> None:
147 """Handle async initialization of the provider; idle models are reloaded on demand."""
148 # Configure the inference runtime before loading any model (see the controller method).
149 self.mass.streams.audio_analysis.ensure_inference_runtime_configured()
150 await self._load_models()
151 self._models_loaded = True
152
153 async def process_pcm_chunk(
154 self,
155 session_id: str,
156 pcm_chunk: bytes,
157 ) -> None:
158 """Process a PCM chunk for beat tracking."""
159 data = self._data.get(session_id)
160 if not data:
161 return
162
163 pcm_mono = await self._run_offloaded(
164 decode_pcm_chunk_to_mono, data.input_audio_format, pcm_chunk
165 )
166 if pcm_mono.size == 0:
167 return
168
169 # Per-chunk VQT for key detection (skip short tail chunks)
170 if len(pcm_mono) >= data.input_audio_format.sample_rate:
171 await self._run_offloaded(
172 self._compute_musical_key_features,
173 pcm_mono,
174 data.input_audio_format.sample_rate,
175 data,
176 )
177
178 data.pcm_buffer.append(pcm_mono)
179 data.pcm_samples += len(pcm_mono)
180
181 # calculate features in 10s blocks to avoid cpu contention
182 if data.pcm_samples >= data.block_samples:
183 await self._process_block(data)
184
185 async def cancel(self, session_id: str) -> None:
186 """Cancel a beat tracking session."""
187 data = self._data.pop(session_id, None)
188 if data:
189 self._clear_session_data(data)
190 await super().cancel(session_id)
191
192 async def _load_models(self) -> None:
193 """Load the Beat This, S-KEY, and FireRed AED models into memory."""
194 (
195 self._beat_this_model,
196 self._beat_this_post_processor,
197 self._skey_vqt,
198 self._skey_chromanet,
199 self._skey_crop,
200 self._spectral_centroid,
201 self._firered_model,
202 self._firered_cmvn_means,
203 self._firered_cmvn_inverse_std,
204 ) = await asyncio.to_thread(self._initialize_models)
205
206 def _free_models(self) -> None:
207 """Release the Beat This, S-KEY, and FireRed AED models."""
208 self._beat_this_model = None
209 self._beat_this_post_processor = None
210 self._skey_vqt = None
211 self._skey_chromanet = None
212 self._skey_crop = None
213 self._spectral_centroid = None
214 self._firered_model = None
215 self._firered_cmvn_means = None
216 self._firered_cmvn_inverse_std = None
217
218 def _initialize_models(
219 self,
220 ) -> tuple[
221 Spect2Frames,
222 DBNDownBeatTracker,
223 VQT,
224 ChromaNet,
225 CropCQT,
226 SpectralCentroid,
227 torch.nn.Module,
228 npt.NDArray[np.float64],
229 npt.NDArray[np.float64],
230 ]:
231 """Initialize ML models (runs in a thread to avoid blocking the event loop)."""
232 beat_this_model = Spect2Frames(checkpoint_path="small0", device=self._device)
233 # torch aarch64 wheels advertise fbgemm in supported_engines but its kernels are x86-only.
234 preference = ("qnnpack", "fbgemm") if is_arm() else ("fbgemm", "qnnpack")
235 supported_engines = torch.backends.quantized.supported_engines
236 quantized_engine = next((e for e in preference if e in supported_engines), None)
237 if quantized_engine is not None and torch.backends.quantized.engine != quantized_engine:
238 torch.backends.quantized.engine = quantized_engine
239 beat_this_model.model = torch.ao.quantization.quantize_dynamic( # type: ignore[no-untyped-call]
240 beat_this_model.model, {torch.nn.Linear}, dtype=torch.qint8
241 )
242 beat_this_post_processor = DBNDownBeatTracker(
243 beats_per_bar=[3, 4], min_bpm=55, max_bpm=215, fps=50
244 )
245 skey_vqt, skey_chromanet, skey_crop = load_skey_components(device=self._device)
246 spectral_centroid = SpectralCentroid(sample_rate=ANALYSIS_SAMPLE_RATE, hop_length=512)
247 firered_model, firered_cmvn_means, firered_cmvn_inverse_std = load_firered_components(
248 device=self._device
249 )
250 return (
251 beat_this_model,
252 beat_this_post_processor,
253 skey_vqt,
254 skey_chromanet,
255 skey_crop,
256 spectral_centroid,
257 firered_model,
258 firered_cmvn_means,
259 firered_cmvn_inverse_std,
260 )
261
262 def _require_loaded[ComponentT](self, component: ComponentT | None, name: str) -> ComponentT:
263 """
264 Return a loaded model component, failing the analysis when it is not resident.
265
266 :param component: The component to resolve; None while the models are unloaded.
267 :param name: Name of the component, used in the failure reason.
268 """
269 if component is None:
270 # The recorder that handles AudioAnalysisError does not log, so warn here or the
271 # unload race leaves nothing behind but a row in the failures table.
272 self.logger.warning("%s is not loaded; analysis will be retried later", name)
273 raise AudioAnalysisError(
274 f"{name} is not loaded",
275 retry_at=utc() + MODEL_FAILURE_RETRY_DELAY,
276 )
277 return component
278
279 async def _start_analysis(
280 self,
281 session_id: str,
282 streamdetails: StreamDetails,
283 audio_format: AudioFormat,
284 ) -> bool:
285 """Start beat tracking analysis for a new track."""
286 if streamdetails.media_type != MediaType.TRACK:
287 # We only want to analyze tracks
288 return False
289
290 block_seconds = 10.0
291
292 needs_resample = audio_format.sample_rate != ANALYSIS_SAMPLE_RATE
293 self._data[session_id] = SmartFadesData(
294 item_id=streamdetails.item_id,
295 provider=streamdetails.provider,
296 input_audio_format=audio_format,
297 block_samples=int(block_seconds * audio_format.sample_rate),
298 features=AdvancedBeatFeatureExtractor(
299 sample_rate=ANALYSIS_SAMPLE_RATE,
300 device=self._device,
301 offload=self._run_offloaded,
302 ),
303 resampler=soxr.ResampleStream(
304 in_rate=audio_format.sample_rate,
305 out_rate=ANALYSIS_SAMPLE_RATE,
306 num_channels=1,
307 dtype="float32",
308 )
309 if needs_resample
310 else None,
311 vocal_resampler=soxr.ResampleStream(
312 in_rate=audio_format.sample_rate,
313 out_rate=FIRERED_SAMPLE_RATE,
314 num_channels=1,
315 dtype="float32",
316 )
317 if audio_format.sample_rate != FIRERED_SAMPLE_RATE
318 else None,
319 vocal_fbank=FireRedFbank(
320 self._require_loaded(self._firered_cmvn_means, "FireRed AED CMVN means"),
321 self._require_loaded(
322 self._firered_cmvn_inverse_std, "FireRed AED CMVN inverse deviations"
323 ),
324 ),
325 )
326 self.logger.debug("Started beat tracking session %s", session_id)
327 return True
328
329 async def _finalize(self, session_id: str) -> AudioAnalysisData | None:
330 """Finalize beat tracking and store results."""
331 data = self._data.pop(session_id, None)
332 if not data:
333 return None
334
335 try:
336 if data.pcm_samples:
337 await self._process_block(data, last=True)
338 else:
339 # The vocal resampler and fbank still need an explicit end-of-input flush.
340 await self._run_offloaded(
341 self._compute_vocal_features,
342 np.empty(0, dtype=np.float32),
343 data,
344 True,
345 )
346
347 final_feats = await data.features.finalize()
348 if final_feats.size:
349 data.beats_feature_blocks.append(final_feats)
350 if not data.beats_feature_blocks:
351 return None
352
353 feats = np.concatenate(data.beats_feature_blocks, axis=0)
354 data.beats_feature_blocks.clear()
355 duration = data.total_pcm_samples / ANALYSIS_SAMPLE_RATE
356
357 all_vqt = None
358 if data.musical_key_feature_blocks:
359 all_vqt = torch.cat(data.musical_key_feature_blocks, dim=-1) # (1, 1, 84, T_total)
360 data.musical_key_feature_blocks.clear()
361
362 if data.vocal_feature_blocks:
363 vocal_features = np.concatenate(data.vocal_feature_blocks)
364 data.vocal_feature_blocks.clear()
365 else:
366 vocal_features = np.empty((0, 80), dtype=np.float32)
367
368 beat_key_result, vocal_activity = await self._run_final_inference(
369 feats,
370 all_vqt,
371 vocal_features,
372 duration,
373 )
374 beats, downbeats, beats_per_bar, key, mode = beat_key_result
375 return self._build_analysis(
376 data,
377 duration,
378 beats,
379 downbeats,
380 beats_per_bar,
381 key,
382 mode,
383 vocal_activity,
384 )
385 finally:
386 self._clear_session_data(data)
387
388 def _build_analysis(
389 self,
390 data: SmartFadesData,
391 duration: float,
392 beats: np.ndarray,
393 downbeats: np.ndarray,
394 beats_per_bar: int,
395 key: str | None,
396 mode: str | None,
397 vocal_activity: np.ndarray,
398 ) -> AudioAnalysisData:
399 """Build the final Smart Fades analysis payload."""
400 bpm = calculate_overall_bpm(beats)
401
402 # mean power per bin: point sampling aliases beat-rate ripple into the bins
403 rms_energy = None
404 energy_peak = 0.0
405 if data.energy_chunks:
406 energy_all = np.concatenate(data.energy_chunks)
407 if len(energy_all) >= 2:
408 rms_energy = aggregate_series_to_bins(energy_all, 1800, power=True)
409 energy_peak = float(rms_energy.max())
410 if energy_peak > 0:
411 rms_energy = rms_energy / energy_peak
412
413 spectral_centroid = None
414 if data.centroid_chunks:
415 centroid_all = np.concatenate(data.centroid_chunks)
416 if len(centroid_all) >= 2:
417 spectral_centroid = aggregate_series_to_bins(centroid_all, 1800)
418 # Zero out centroid where energy is negligible (noise dominates)
419 if rms_energy is not None:
420 spectral_centroid[rms_energy < 0.01] = 0.0
421
422 vocal_activity_bins = (
423 aggregate_series_to_bins(vocal_activity, 1800)
424 if vocal_activity.size
425 else np.zeros(1800, dtype=np.float32)
426 )
427 extra_data: dict[str, Any] = {"vocal_activity": vocal_activity_bins.tolist()}
428 if energy_peak > 0 and data.frequency_band_chunks:
429 extra_data["band_rms"] = {
430 name: (
431 aggregate_series_to_bins(np.concatenate(chunks), 1800, power=True) / energy_peak
432 ).tolist()
433 for name, chunks in data.frequency_band_chunks.items()
434 }
435
436 analysis = AudioAnalysisData(
437 bpm=bpm,
438 # the model stores plain float lists (numpy-free); convert the analysis arrays
439 beats=beats.tolist(),
440 downbeats=downbeats.tolist(),
441 duration=duration,
442 rms_energy=rms_energy.tolist() if rms_energy is not None else None,
443 spectral_centroid=(
444 spectral_centroid.tolist() if spectral_centroid is not None else None
445 ),
446 key=key,
447 mode=mode,
448 extra_data=extra_data,
449 beats_per_bar=beats_per_bar or None,
450 )
451 self.logger.debug(
452 "Beat analysis for %s: BPM=%.1f, %d beats, %d downbeats, key=%s",
453 data.item_id,
454 bpm,
455 len(beats),
456 len(downbeats),
457 f"{key} {mode}" if key else "unknown",
458 )
459 return analysis
460
461 async def _process_block(self, data: SmartFadesData, *, last: bool = False) -> None:
462 """Resample accumulated PCM buffer and extract features."""
463 start_time = time.perf_counter()
464 pcm_raw = (
465 np.concatenate(data.pcm_buffer) if data.pcm_buffer else np.empty(0, dtype=np.float32)
466 )
467 data.pcm_buffer.clear()
468 data.pcm_samples = 0
469
470 if data.resampler is not None:
471 pcm_22k = await self._run_offloaded(data.resampler.resample_chunk, pcm_raw, last)
472 else:
473 pcm_22k = pcm_raw
474
475 data.total_pcm_samples += len(pcm_22k)
476
477 if pcm_22k.size:
478 feats, _, _ = await asyncio.gather(
479 data.features.process_pcm(pcm_22k),
480 self._run_offloaded(
481 self._compute_energy_and_spectral_centroids,
482 pcm_22k,
483 data,
484 ),
485 self._run_offloaded(self._compute_vocal_features, pcm_raw, data, last),
486 )
487 else:
488 await self._run_offloaded(self._compute_vocal_features, pcm_raw, data, last)
489 feats = np.empty((0, 128), dtype=np.float32)
490
491 if feats.size:
492 data.beats_feature_blocks.append(feats)
493
494 elapsed_ms = (time.perf_counter() - start_time) * 1000
495 self.logger.log(VERBOSE_LOG_LEVEL, "Processed 10s of PCM chunks in %.1fms", elapsed_ms)
496
497 def _compute_energy_and_spectral_centroids(
498 self, pcm_22k: np.ndarray, data: SmartFadesData
499 ) -> None:
500 """Compute fine-resolution RMS energy and spectral centroid for a block."""
501 sr = ANALYSIS_SAMPLE_RATE
502 # RMS energy in 100ms windows, including partial final window
503 window_samples = sr // 10 # 2205 samples = 100ms
504 if len(pcm_22k) > 0:
505 n_full = len(pcm_22k) // window_samples
506 rms_list = []
507 if n_full > 0:
508 frames = pcm_22k[: n_full * window_samples].reshape(n_full, window_samples)
509 rms_list.append(np.sqrt(np.mean(frames**2, axis=1)))
510 remainder = len(pcm_22k) - n_full * window_samples
511 if remainder > 0:
512 tail = pcm_22k[n_full * window_samples :]
513 rms_list.append(np.array([np.sqrt(np.mean(tail**2))]))
514 if rms_list:
515 data.energy_chunks.append(np.concatenate(rms_list).astype(np.float32))
516
517 band_frames = compute_band_rms_frames(pcm_22k, sr, window_samples)
518 for name, frames in band_frames.items():
519 data.frequency_band_chunks.setdefault(name, []).append(frames)
520
521 # Spectral centroid: keep per-frame (hop_length=512, ~43 frames/s)
522 # Skip short tail buffers: STFT reflect-pad requires len > n_fft // 2.
523 spectral_centroid = self._require_loaded(
524 self._spectral_centroid, "SpectralCentroid transform"
525 )
526 if len(pcm_22k) >= spectral_centroid.n_fft:
527 pcm_tensor = torch.from_numpy(pcm_22k)
528 centroid_frames = spectral_centroid(pcm_tensor.unsqueeze(0)).squeeze(0).numpy()
529 # digitally-silent frames divide 0/0 into NaN; treat them as 0 Hz like
530 # other negligible-energy frames so no non-finite value is ever stored
531 np.nan_to_num(centroid_frames, copy=False, nan=0.0, posinf=0.0, neginf=0.0)
532 if len(centroid_frames) > 0:
533 data.centroid_chunks.append(centroid_frames.astype(np.float32))
534
535 def _compute_vocal_features(
536 self,
537 pcm_raw: np.ndarray,
538 data: SmartFadesData,
539 last: bool,
540 ) -> None:
541 """Resample source PCM to 16 kHz and extract FireRed fbank features."""
542 fbank = data.vocal_fbank
543 resampler = data.vocal_resampler
544 if fbank is None:
545 return
546 pcm_16k = resampler.resample_chunk(pcm_raw, last) if resampler is not None else pcm_raw
547 features = fbank.process(pcm_16k)
548 if features.size:
549 data.vocal_feature_blocks.append(features)
550 if last:
551 final_features = fbank.finalize()
552 if final_features.size:
553 data.vocal_feature_blocks.append(final_features)
554
555 async def _run_final_inference(
556 self,
557 beat_features: np.ndarray,
558 key_features: torch.Tensor | None,
559 vocal_features: np.ndarray,
560 duration: float,
561 ) -> tuple[tuple[np.ndarray, np.ndarray, int, str | None, str | None], np.ndarray]:
562 """Run beat/key and vocal inference branches; the first failure cancels the other."""
563 try:
564 async with asyncio.TaskGroup() as task_group:
565 beat_key_task = task_group.create_task(
566 self._infer_beats_and_key(beat_features, key_features)
567 )
568 vocal_task = task_group.create_task(
569 self._infer_vocal_activity(vocal_features, duration)
570 )
571 except ExceptionGroup as group:
572 # Unwrap for the base class; prefer the beat/key error since it decides
573 # permanent vs retryable failure recording.
574 beat_key_error = None if beat_key_task.cancelled() else beat_key_task.exception()
575 primary = beat_key_error or group.exceptions[0]
576 for error in group.exceptions:
577 if error is not primary:
578 self.logger.debug("FireRed vocal inference also failed: %s", error)
579 raise primary from primary.__cause__
580 return beat_key_task.result(), vocal_task.result()
581
582 async def _infer_beats_and_key(
583 self,
584 beat_features: np.ndarray,
585 key_features: torch.Tensor | None,
586 ) -> tuple[np.ndarray, np.ndarray, int, str | None, str | None]:
587 """Run beat inference followed by musical key inference."""
588 # Resolved before the beat stage: an idle unload may clear the field while it runs.
589 chromanet = self._require_loaded(self._skey_chromanet, "S-KEY ChromaNet")
590 beats, downbeats, beats_per_bar = await self._infer_beat_timings(beat_features)
591 if len(beats) < 2:
592 raise AudioAnalysisError("no rhythmic beat detected")
593 key, mode = await self._run_offloaded(self._infer_musical_key, chromanet, key_features)
594 return beats, downbeats, beats_per_bar, key, mode
595
596 async def _infer_vocal_activity(
597 self,
598 features: np.ndarray,
599 duration: float,
600 ) -> np.ndarray:
601 """Run FireRed AED inference and return the 100 ms vocal timeline."""
602 model = self._require_loaded(self._firered_model, "FireRed AED model")
603 try:
604 compute_seconds = 0.0
605 chunks = []
606 for chunk, core_offset, core_length in split_firered_features(features):
607 chunk_probabilities, elapsed = await self._run_offloaded_timed(
608 infer_firered_chunk,
609 model,
610 chunk,
611 self._device,
612 )
613 compute_seconds += elapsed
614 chunks.append(chunk_probabilities[core_offset : core_offset + core_length])
615 frame_probabilities = (
616 np.concatenate(chunks) if chunks else np.empty((0, 3), dtype=np.float32)
617 )
618 probabilities = vocal_activity_probabilities(frame_probabilities, duration)
619 except asyncio.CancelledError:
620 raise
621 except Exception as err:
622 # Avoid permanently suppressing beat and key results for transient model failures.
623 raise AudioAnalysisError(
624 f"FireRed vocal inference failed: {err}",
625 retry_at=utc() + MODEL_FAILURE_RETRY_DELAY,
626 ) from err
627 self.logger.log(
628 VERBOSE_LOG_LEVEL,
629 "FireRed vocal inference: %.1fms compute over %d frames",
630 compute_seconds * 1000,
631 len(features),
632 )
633 return probabilities
634
635 def _compute_musical_key_features(
636 self, pcm_mono: np.ndarray, sample_rate: int, data: SmartFadesData
637 ) -> None:
638 """Extract VQT features for S-KEY key detection."""
639 vqt = self._require_loaded(self._skey_vqt, "S-KEY VQT")
640 crop = self._require_loaded(self._skey_crop, "S-KEY CropCQT")
641 if sample_rate != ANALYSIS_SAMPLE_RATE:
642 pcm_mono = soxr.resample(pcm_mono, sample_rate, ANALYSIS_SAMPLE_RATE)
643 pcm_tensor = torch.from_numpy(pcm_mono)
644 with torch.inference_mode():
645 vqt_input = pcm_tensor.unsqueeze(0).unsqueeze(0) # (1, 1, samples)
646 vqt_out = vqt(vqt_input) # (1, 1, n_bins, T)
647 cropped = crop(vqt_out, torch.zeros(1)) # (1, 1, 84, T)
648 data.musical_key_feature_blocks.append(cropped.cpu())
649
650 def _infer_musical_key(
651 self, chromanet: torch.nn.Module, vqt_features: torch.Tensor | None
652 ) -> tuple[str | None, str | None]:
653 """
654 Run S-KEY ChromaNet inference to detect musical key.
655
656 :param chromanet: The ChromaNet module to run.
657 :param vqt_features: Accumulated VQT features, or None when the track had too few.
658 """
659 if vqt_features is None or vqt_features.shape[-1] < 128:
660 return None, None
661 start = time.perf_counter()
662 with torch.no_grad():
663 logits = chromanet(vqt_features.to(self._device))
664 key_idx = int(logits.argmax(dim=-1).item())
665 key_name = SKEY_KEY_MAP[key_idx] # e.g. "C# Major"
666 parts = key_name.split()
667 self.logger.log(
668 VERBOSE_LOG_LEVEL,
669 "ChromaNet key inference: %.1fms, detected key=%s %s",
670 (time.perf_counter() - start) * 1000,
671 parts[0],
672 parts[1],
673 )
674 return parts[0], parts[1].lower()
675
676 async def _infer_beat_timings(self, feats: np.ndarray) -> tuple[np.ndarray, np.ndarray, int]:
677 """
678 Run Beat This model inference to detect beat/downbeat timings and the meter.
679
680 :param feats: Log-mel features for the whole track, shaped (frames, mel bins).
681 """
682 # Resolved once and passed down: an idle unload may clear the fields while the
683 # windows below are still being dispatched.
684 beat_this = self._require_loaded(self._beat_this_model, "Beat This model")
685 post_processor = self._require_loaded(
686 self._beat_this_post_processor, "Beat This post-processor"
687 )
688
689 spect = torch.from_numpy(feats).to(self._device)
690 windows, starts = split_piece(
691 spect,
692 BEAT_WINDOW_FRAMES,
693 border_size=BEAT_WINDOW_BORDER_FRAMES,
694 avoid_short_end=True,
695 )
696 predictions = []
697 model_seconds = 0.0
698 for window in windows:
699 prediction, elapsed = await self._run_offloaded_timed(
700 self._infer_beat_window, beat_this.model, window
701 )
702 predictions.append(prediction)
703 model_seconds += elapsed
704 await self._pace_beat_windows(elapsed)
705
706 (beats, downbeats, beats_per_bar), post_seconds = await self._run_offloaded_timed(
707 self._decode_beat_timings, post_processor, predictions, starts, len(spect)
708 )
709 self.logger.log(
710 VERBOSE_LOG_LEVEL,
711 "Model inference: %.1fms compute over %d windows, postprocessing: %.1fms, "
712 "detected %d beats, %d downbeats",
713 model_seconds * 1000,
714 len(windows),
715 post_seconds * 1000,
716 len(beats),
717 len(downbeats),
718 )
719 return beats, downbeats, beats_per_bar
720
721 async def _pace_beat_windows(self, window_seconds: float) -> None:
722 """
723 Idle for as long as the beat inference window that just finished took to compute.
724
725 Only while a player streams; idle and background analysis run at full speed.
726
727 :param window_seconds: Compute time of the window that just finished.
728 """
729 if not self.mass.streams.audio_analysis.playback_active():
730 return
731 await asyncio.sleep(window_seconds * BEAT_WINDOW_PACE_RATIO)
732
733 @staticmethod
734 def _infer_beat_window(model: torch.nn.Module, window: torch.Tensor) -> dict[str, torch.Tensor]:
735 """
736 Run one Beat This window and return its beat and downbeat logits.
737
738 :param model: The Beat This module to run.
739 :param window: One window of log-mel features, shaped (frames, mel bins).
740 """
741 # inference_mode is thread-local, so it has to be entered on the worker thread.
742 with torch.inference_mode():
743 prediction = model(window.unsqueeze(0))
744 return {"beat": prediction["beat"][0], "downbeat": prediction["downbeat"][0]}
745
746 def _decode_beat_timings(
747 self,
748 post_processor: DBNDownBeatTracker,
749 predictions: list[dict[str, torch.Tensor]],
750 starts: np.ndarray,
751 total_frames: int,
752 ) -> tuple[np.ndarray, np.ndarray, int]:
753 """
754 Stitch per-window logits back together and decode them into beat timings.
755
756 :param post_processor: The DBN decoder to run on the stitched activations.
757 :param predictions: Per-window beat/downbeat logits, in window order.
758 :param starts: Frame offset of each window, as returned by split_piece.
759 :param total_frames: Frame count of the whole track.
760 """
761 with torch.inference_mode():
762 beat_logits, downbeat_logits = aggregate_prediction(
763 predictions,
764 starts,
765 total_frames,
766 BEAT_WINDOW_FRAMES,
767 BEAT_WINDOW_BORDER_FRAMES,
768 BEAT_WINDOW_OVERLAP_MODE,
769 self._device,
770 )
771 dbn_out, beats_per_bar = post_processor(
772 self._beat_activations(beat_logits.float(), downbeat_logits.float())
773 )
774 beats = dbn_out[:, 0]
775 downbeats = dbn_out[dbn_out[:, 1] == 1, 0]
776 return beats, downbeats, beats_per_bar
777
778 @staticmethod
779 def _beat_activations(beat_logits: torch.Tensor, downbeat_logits: torch.Tensor) -> np.ndarray:
780 """
781 Convert beat/downbeat logits into the (T, 2) activations the DBN expects.
782
783 :param beat_logits: Per-frame beat logits for the whole track.
784 :param downbeat_logits: Per-frame downbeat logits for the whole track.
785 """
786 beat_prob = torch.sigmoid(beat_logits).cpu().numpy()
787 downbeat_prob = torch.sigmoid(downbeat_logits).cpu().numpy()
788 epsilon = 1e-5
789 beat_prob = beat_prob * (1 - epsilon) + epsilon / 2
790 downbeat_prob = downbeat_prob * (1 - epsilon) + epsilon / 2
791 return np.column_stack(
792 [
793 np.maximum(beat_prob - downbeat_prob, epsilon / 2),
794 downbeat_prob,
795 ]
796 )
797
798 def _clear_session_data(self, data: SmartFadesData) -> None:
799 """Release all state retained for an analysis session."""
800 data.pcm_buffer.clear()
801 data.beats_feature_blocks.clear()
802 data.energy_chunks.clear()
803 data.centroid_chunks.clear()
804 data.frequency_band_chunks.clear()
805 data.musical_key_feature_blocks.clear()
806 data.vocal_feature_blocks.clear()
807 data.resampler = None
808 data.vocal_resampler = None
809 data.vocal_fbank = None
810