trackembeddingapi / tests /conftest.py
kxmWebwe's picture
update
330f477
Raw
History Blame Contribute Delete
4.32 kB
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
@pytest.fixture
def service_bundle():
return make_service()