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