Spaces:
Running
Running
| from __future__ import annotations | |
| import hashlib | |
| from pathlib import Path | |
| import numpy as np | |
| import pytest | |
| from openmusic_analysis.analyzers.clap import ( | |
| ClapGlobalAudioAnalyzer, | |
| ClapTemporalAudioAnalyzer, | |
| ) | |
| from openmusic_analysis.analyzers.lyrics import BGEM3LyricsAnalyzer, LyricsPreprocessor | |
| from openmusic_analysis.application import MusicAnalysisService | |
| from openmusic_analysis.audio.decoder import AudioMetadata, DecodedAudio | |
| from openmusic_analysis.errors import AudioDecodeError | |
| from openmusic_analysis.registry import ModelRegistry | |
| from openmusic_analysis.settings import GlobalAudioConfig, LyricsConfig, TemporalAudioConfig | |
| class FakeAudioEncoder: | |
| dimension = 4 | |
| loaded = True | |
| def __init__(self) -> None: | |
| self.calls: list[list[np.ndarray]] = [] | |
| async def ready(self) -> None: | |
| return None | |
| async def encode(self, windows: list[np.ndarray]) -> np.ndarray: | |
| self.calls.append(windows) | |
| rows = [] | |
| for window in windows: | |
| rows.append( | |
| [ | |
| float(np.mean(window)), | |
| float(np.std(window)), | |
| float(window[0]), | |
| float(window[-1]) + 1.0, | |
| ] | |
| ) | |
| return np.asarray(rows, dtype=np.float32) | |
| class FailingAudioEncoder(FakeAudioEncoder): | |
| async def encode(self, windows: list[np.ndarray]) -> np.ndarray: | |
| raise RuntimeError("internal model detail must not leak") | |
| class FakeTextEncoder: | |
| dimension = 6 | |
| loaded = True | |
| async def ready(self) -> None: | |
| return None | |
| def count_tokens(self, text: str) -> int: | |
| return len(text.split()) + 2 | |
| def split_tokens(self, text: str, max_tokens: int) -> list[str]: | |
| words = text.split() | |
| size = max(1, max_tokens - 2) | |
| return [" ".join(words[index : index + size]) for index in range(0, len(words), size)] | |
| async def encode(self, texts: list[str], batch_size: int) -> np.ndarray: | |
| rows = [] | |
| for text in texts: | |
| digest = hashlib.sha256(text.encode("utf-8")).digest() | |
| rows.append([float(value + 1) for value in digest[: self.dimension]]) | |
| return np.asarray(rows, dtype=np.float32) | |
| class FakeDecoder: | |
| canonical_sample_rate = 10 | |
| def __init__(self, waveform: np.ndarray | None = None) -> None: | |
| self.waveform = ( | |
| np.asarray(waveform, dtype=np.float32) | |
| if waveform is not None | |
| else np.linspace(-1.0, 1.0, 95, dtype=np.float32) | |
| ) | |
| self.calls = 0 | |
| def decode(self, source_path: str | Path) -> DecodedAudio: | |
| self.calls += 1 | |
| if Path(source_path).read_bytes().startswith(b"bad"): | |
| raise AudioDecodeError() | |
| return DecodedAudio( | |
| waveform=self.waveform, | |
| sample_rate=self.canonical_sample_rate, | |
| metadata=AudioMetadata( | |
| duration_ms=int(round(self.waveform.size * 1000 / self.canonical_sample_rate)), | |
| source_sample_rate=self.canonical_sample_rate, | |
| source_channels=1, | |
| source_format="fake", | |
| ), | |
| ) | |
| def make_service( | |
| *, | |
| waveform: np.ndarray | None = None, | |
| audio_encoder: FakeAudioEncoder | None = None, | |
| ) -> tuple[MusicAnalysisService, FakeDecoder, FakeAudioEncoder]: | |
| audio_encoder = audio_encoder or FakeAudioEncoder() | |
| decoder = FakeDecoder(waveform) | |
| global_config = GlobalAudioConfig( | |
| sample_rate=10, | |
| window_seconds=2.0, | |
| target_windows=4, | |
| minimum_audio_seconds=1.0, | |
| ) | |
| temporal_config = TemporalAudioConfig( | |
| sample_rate=10, | |
| window_seconds=2.0, | |
| hop_seconds=2.0, | |
| max_segments=4, | |
| minimum_audio_seconds=1.0, | |
| ) | |
| analyzers = [ | |
| ClapGlobalAudioAnalyzer(audio_encoder, global_config), | |
| ClapTemporalAudioAnalyzer(audio_encoder, temporal_config), | |
| BGEM3LyricsAnalyzer( | |
| FakeTextEncoder(), | |
| LyricsPreprocessor(), | |
| LyricsConfig(max_chunk_tokens=12, batch_size=4), | |
| ), | |
| ] | |
| service = MusicAnalysisService( | |
| registry=ModelRegistry(analyzers), decoder=decoder, device="cpu" | |
| ) | |
| return service, decoder, audio_encoder | |
| def service_bundle(): | |
| return make_service() | |