sesa-gpu / tests /test_batch_duration.py
John6666's picture
Upload 44 files
81ba775 verified
Raw
History Blame Contribute Delete
2.56 kB
from pathlib import Path
from src import duration, ui
def _estimate(monkeypatch, tmp_path, *, files=2, package_models=None, safe=False):
paths = []
for index in range(files):
path = tmp_path / f"item-{index}.wav"
path.write_bytes(b"x")
paths.append(path)
monkeypatch.setattr(duration, "probe_duration", lambda path: 15.0)
return duration.calculate_duration_estimate(
paths,
[],
package_models or ["UVR_MDXNET_KARA_2.onnx"],
"None",
"avg_wave",
"OPUS",
0,
0,
duration.DURATION_MODE_SEMI_AUTO,
180,
12,
1.25,
safe,
)
def test_two_item_onnx_batch_declares_cold_start_aware_duration(monkeypatch, tmp_path):
estimate = _estimate(monkeypatch, tmp_path)
assert estimate.seconds == 78
assert estimate.first_inference_warmup_seconds == duration.ONNX_FIRST_INFERENCE_WARMUP_SECONDS
assert estimate.additional_batch_item_overhead_seconds == duration.ADDITIONAL_BATCH_ITEM_DISPATCH_SECONDS
assert estimate.calibration_revision == duration.DURATION_CALIBRATION_REVISION
assert estimate.onnx_model_count == 1
def test_safe_mode_applies_after_cold_start_and_batch_overhead(monkeypatch, tmp_path):
base = _estimate(monkeypatch, tmp_path, safe=False)
safe = _estimate(monkeypatch, tmp_path, safe=True)
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 * duration.SAFE_MODE_MULTIPLIER)
def test_non_onnx_profile_uses_non_onnx_warmup(monkeypatch, tmp_path):
estimate = _estimate(
monkeypatch,
tmp_path,
files=1,
package_models=["UVR-De-Reverb-aufr33-jarredou.pth"],
)
assert estimate.first_inference_warmup_seconds == duration.NON_ONNX_FIRST_INFERENCE_WARMUP_SECONDS
assert estimate.onnx_model_count == 0
assert estimate.non_onnx_model_count == 1
def test_ui_wrapper_passes_explicit_batch_failure_policy(monkeypatch):
captured = {}
def fake_prepare(*args, **kwargs):
captured["kwargs"] = kwargs
return "state", "markdown", "preflight.json", "prepare.log"
monkeypatch.setattr(ui, "prepare_job", fake_prepare)
values = list(range(34)) + [False]
result = ui.prepare_separation_with_progress(*values, progress=None)
assert result[0] == "state"
assert captured["kwargs"]["batch_continue_on_item_error"] is False