Spaces:
Running
Running
| from __future__ import annotations | |
| import numpy as np | |
| import pytest | |
| from openmusic_analysis.analyzers.lyrics import BGEM3LyricsAnalyzer, LyricsPreprocessor | |
| from openmusic_analysis.errors import AnalysisError | |
| from openmusic_analysis.settings import LyricsConfig | |
| from .conftest import FakeTextEncoder | |
| def analyzer(max_tokens: int = 12) -> BGEM3LyricsAnalyzer: | |
| return BGEM3LyricsAnalyzer( | |
| FakeTextEncoder(), | |
| LyricsPreprocessor(), | |
| LyricsConfig(max_chunk_tokens=max_tokens, batch_size=4), | |
| ) | |
| async def test_english_and_russian_are_finite_and_normalized(lyrics): | |
| result = await analyzer().analyze(lyrics) | |
| vector = np.asarray(result.embedding) | |
| assert result.dimension == 6 | |
| assert np.isfinite(vector).all() | |
| assert np.linalg.norm(vector) == pytest.approx(1.0, abs=1e-6) | |
| def test_preprocessor_preserves_multiline_sections_punctuation_and_repeated_chorus(): | |
| lyrics = "[ar:metadata]\r\n[Verse]\r\nHello, world!\r\n\r\n[Chorus]\r\nAgain!\r\n[Chorus]\r\nAgain!" | |
| prepared = LyricsPreprocessor().prepare(lyrics) | |
| assert "[ar:metadata]" not in prepared.normalized_text | |
| assert "Hello, world!" in prepared.normalized_text | |
| assert [section.label.lower() for section in prepared.sections] == [ | |
| "verse", | |
| "chorus", | |
| "chorus", | |
| ] | |
| assert [section.text for section in prepared.sections].count("Again!") == 2 | |
| async def test_long_lyrics_use_token_chunks_and_deterministic_aggregation(): | |
| lyrics = "[Verse]\n" + " ".join(f"word{index}" for index in range(80)) | |
| first = await analyzer(max_tokens=10).analyze(lyrics) | |
| second = await analyzer(max_tokens=10).analyze(lyrics) | |
| assert first.analysis["chunk_count"] > 1 | |
| assert max(first.analysis["chunk_token_counts"]) <= 10 | |
| assert first.embedding == pytest.approx(second.embedding, abs=1e-7) | |
| async def test_empty_lyrics_are_rejected(): | |
| with pytest.raises(AnalysisError) as raised: | |
| await analyzer().analyze(" \n\t ") | |
| assert raised.value.code == "INVALID_LYRICS" | |