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