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