sesa-gpu / tests /test_duration.py
John6666's picture
Upload 44 files
81ba775 verified
Raw
History Blame Contribute Delete
3.38 kB
from pathlib import Path
from src import duration
def test_manual_duration_is_fixed_and_clamped(monkeypatch, tmp_path):
sample = tmp_path / "sample.wav"
sample.write_bytes(b"x")
monkeypatch.setattr(duration, "probe_duration", lambda path: 120.0)
estimate = duration.calculate_duration_estimate(
[sample], [], [], "None", "avg_wave", "FLAC", 0, 0,
duration.DURATION_MODE_MANUAL, 999, 12, 1.25,
)
assert estimate.seconds == duration.ZERO_GPU_MAX_SECONDS
def test_semi_auto_increases_for_more_models_and_parameters(monkeypatch, tmp_path):
sample = tmp_path / "sample.wav"
sample.write_bytes(b"x")
monkeypatch.setattr(duration, "probe_duration", lambda path: 180.0)
base = duration.calculate_duration_estimate(
[sample], [], ["UVR-MDX-NET-Inst_HQ_5.onnx"], "None",
"avg_wave", "FLAC", 0, 0,
duration.DURATION_MODE_SEMI_AUTO, 180, 8, 1.0,
)
heavier = duration.calculate_duration_estimate(
[sample], [], ["UVR-MDX-NET-Inst_HQ_5.onnx", "UVR-De-Reverb-aufr33-jarredou.pth"], "None",
"avg_fft", "MP3", 2, 120,
duration.DURATION_MODE_SEMI_AUTO, 180, 8, 1.0,
)
assert heavier.seconds > base.seconds
assert heavier.model_count == 2
def test_duration_callable_uses_app_argument_positions(monkeypatch, tmp_path):
sample = tmp_path / "sample.wav"
sample.write_bytes(b"x")
monkeypatch.setattr(duration, "probe_duration", lambda path: 60.0)
args = [
[sample], [], [], "None", "", "", "", "main", "", "", "",
"avg_wave", "FLAC", "All stems", 0, 0,
duration.DURATION_MODE_MANUAL, 75, 12, 1.25,
]
assert duration.estimate_gpu_seconds(*args) == 75
def test_safe_mode_defaults_off_and_adds_exact_30_percent(monkeypatch, tmp_path):
sample = tmp_path / "sample.wav"
sample.write_bytes(b"x")
monkeypatch.setattr(duration, "probe_duration", lambda path: 120.0)
base = duration.calculate_duration_estimate(
[sample], [], [], "None", "avg_wave", "FLAC", 0, 0,
duration.DURATION_MODE_SEMI_AUTO, 180, 10, 1.0, False,
)
safe = duration.calculate_duration_estimate(
[sample], [], [], "None", "avg_wave", "FLAC", 0, 0,
duration.DURATION_MODE_SEMI_AUTO, 180, 10, 1.0, True,
)
assert base.safe_mode_requested is False
assert base.safe_mode_applied is False
assert base.safe_mode_multiplier == 1.0
assert safe.safe_mode_requested is True
assert safe.safe_mode_applied is True
assert safe.safe_mode_multiplier == duration.SAFE_MODE_MULTIPLIER
assert safe.base_unclamped_seconds == base.base_unclamped_seconds
assert safe.unclamped_seconds == base.unclamped_seconds * duration.SAFE_MODE_MULTIPLIER
assert safe.seconds == duration._clamp_seconds(base.unclamped_seconds * 1.30)
def test_safe_mode_never_changes_manual_duration(monkeypatch, tmp_path):
sample = tmp_path / "sample.wav"
sample.write_bytes(b"x")
monkeypatch.setattr(duration, "probe_duration", lambda path: 120.0)
estimate = duration.calculate_duration_estimate(
[sample], [], [], "None", "avg_wave", "FLAC", 0, 0,
duration.DURATION_MODE_MANUAL, 75, 10, 1.0, True,
)
assert estimate.seconds == 75
assert estimate.safe_mode_requested is True
assert estimate.safe_mode_applied is False
assert estimate.safe_mode_multiplier == 1.0