music3lab / tests /test_long_reference_style.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
2.89 kB
from __future__ import annotations
import numpy as np
from music3lab import long_reference_style as lrs
def test_balanced_reference_style_bestofn_preserves_raws_and_selects_eligible(tmp_path) -> None:
"""Contract for the long-form bridge; all callbacks keep this CPU-only."""
policy = lrs.resolve_policy("balanced")
assert (policy.candidates, policy.duration_seconds, policy.frames, policy.chunk_count) == (4, 60.0, 1500, 15)
source = {
"tempo_bpm": 126.0, "key": "D", "mode": "minor",
"energy": 0.42, "stereo_width": 0.25,
}
first = lrs.build_candidate_plan(source, root_seed=17, index=0, policy=policy)
again = lrs.build_candidate_plan(source, root_seed=17, index=0, policy=policy)
other = lrs.build_candidate_plan(source, root_seed=17, index=1, policy=policy)
assert first == again and first["plan_seed"] != other["plan_seed"]
assert first["plan_seed"] != first["ar_flow_seed"]
assert first["executed_seed_fields"] == ["ar_flow_seed"]
assert first["independent_ar_flow_seed_supported"] is False
assert {"tempo_bpm", "key", "mode", "energy", "stereo_width", "texture", "arrangement"} <= set(first)
assert first["instrumentation_measured"] is False
calls: list[dict[str, object]] = []
def render(**kwargs: object) -> np.ndarray:
calls.append(kwargs)
# Candidate 1 is early EOS; candidate 2 is a genuine decoder overrun.
length = {0: 1200, 1: 800, 2: 1230, 3: 1200}[int(kwargs["index"])]
t = np.arange(length, dtype=np.float32) / 20.0
return np.stack((0.08 * np.sin(2 * np.pi * (2 + int(kwargs["index"])) * t),
0.07 * np.sin(2 * np.pi * (3 + int(kwargs["index"])) * t)))
def profile(audio: np.ndarray) -> dict[str, float]:
return {"distance": float(abs(audio.shape[-1] - 1200) / 1200), "energy": 0.42}
result = lrs.generate_long_reference_style(
render=render, source_analysis=source, source_profile={"distance": 0.0},
policy="balanced", root_seed=17, output=tmp_path / "selected.wav",
sample_rate=20, profile=profile,
negative="no clipping, no source copy, no tempo drift, no vocals",
)
assert len(calls) == 4 and result.output.exists() and result.selected.eligible
assert all((result.candidate_root / f"candidate-{index:03d}.raw.wav").exists() for index in range(4))
assert not (result.candidate_root / "candidate-001.trimmed.wav").exists()
assert (result.candidate_root / "candidate-002.trimmed.wav").exists()
assert result.candidates[1].eligible is False and "early_eos" in result.candidates[1].rejections
assert result.negative_constraints["vocals"]["status"] == "NOT_ENFORCEABLE"
assert result.negative_constraints["native_negative_prompt_used"] is False
assert result.capability == "long_reference_internal_text_bridge_plus_postrender_ranking"