sesa-gpu / tests /test_batch_orchestration.py
John6666's picture
Upload 44 files
81ba775 verified
Raw
History Blame Contribute Delete
11.2 kB
import json
import zipfile
from pathlib import Path
from src import service
from src.jobs import JobPaths
def _job(tmp_path: Path) -> JobPaths:
root = tmp_path / "job"
job = JobPaths("a" * 32, root, root / "input", root / "work", root / "output")
for path in (job.input_dir, job.work_dir, job.output_dir, job.logs_dir, job.bundle_dir, job.config_dir):
path.mkdir(parents=True, exist_ok=True)
return job
def _config(job: JobPaths, prepared: list[Path]) -> dict:
return {
"schema": "sesa-preparation-v21",
"reproducibility_schema": "sesa-reproducibility-v21",
"status": "PREPARED",
"job_id": job.job_id,
"preparation_wall_seconds": 1.0,
"input_total_seconds": 30.0,
"inputs": [
{
"display_name": f"source-{index}.wav",
"prepared_path": path.relative_to(job.root).as_posix(),
"prepared_sha256": f"sha-{index}",
"prepared_duration_seconds": 15.0,
}
for index, path in enumerate(prepared)
],
"batch": {
"schema": "sesa-batch-plan-v1",
"enabled": True,
"input_count": 2,
"execution_order": "sequential-inputs-shared-loaded-model",
"model_load_scope": "once-per-job",
"continue_on_item_error": True,
"output_mapping": "explicit-per-separate-call",
"status": "PREPARED",
},
"models": {
"selected_filenames": ["model.ckpt"],
"custom_records": [],
"ensemble_algorithm": "avg_wave",
},
"output": {
"format": "FLAC",
"bitrate": "Auto",
"sample_rate": 44100,
"normalization_threshold": 0.9,
"amplification_threshold": 0.0,
"single_stem": "All stems",
"bundle_layout": "Group by input",
},
"parameters": {
"pitch_shift": 0,
"chunk_duration": 0,
"allow_cpu_fallback": False,
"log_level": "INFO",
},
"duration": {"seconds": 180},
}
def _patch_runtime(monkeypatch, job: JobPaths, config: dict, fake_separator):
monkeypatch.setattr(service, "parse_prepared_state", lambda value: (job, config))
monkeypatch.setattr(service, "_custom_records_from_config", lambda value: [])
monkeypatch.setattr(service.torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(service.torch.cuda, "get_device_name", lambda index: "Fake CUDA")
monkeypatch.setattr(service, "create_separator", lambda **kwargs: fake_separator)
monkeypatch.setattr(service, "release_accelerators", lambda separator=None: None)
def timed(function, use_cuda):
return function(), 0.1, None
monkeypatch.setattr(service, "synchronized_wall_time", timed)
def test_batch_runs_each_item_sequentially_with_one_model_load(monkeypatch, tmp_path):
job = _job(tmp_path)
prepared = []
for index in range(2):
path = job.work_dir / f"{index:02d}_mix.wav"
path.write_bytes(f"input-{index}".encode())
prepared.append(path)
config = _config(job, prepared)
class FakeSeparator:
def __init__(self):
self.load_calls = 0
self.separate_calls = []
def load_model(self, value):
self.load_calls += 1
def separate(self, values):
assert len(values) == 1
source = Path(values[0])
self.separate_calls.append(source.name)
output = job.output_dir / f"{source.stem}_(Vocals).flac"
output.write_bytes(source.name.encode())
return [str(output)]
separator = FakeSeparator()
_patch_runtime(monkeypatch, job, config, separator)
result = service.run_prepared_job("state")
status, _, _, outputs, archive, _, _, repro_path = result
assert status.startswith("### Completed")
assert separator.load_calls == 1
assert separator.separate_calls == [path.name for path in prepared]
assert len(outputs) == 3
stem_outputs = [Path(value) for value in outputs if Path(value).parent == job.output_dir]
assert len(stem_outputs) == 2
assert Path(outputs[-1]).name == "batch_manifest.json"
repro = json.loads(Path(repro_path).read_text(encoding="utf-8"))
batch = repro["execution"]["batch"]
assert batch["status"] == "COMPLETED"
assert batch["completed_count"] == 2
assert batch["failed_count"] == 0
assert batch["model_load_count"] == 1
assert sorted(item["input_index"] for item in repro["outputs"]) == [0, 1]
assert all("_input_" in path.name for path in stem_outputs)
with zipfile.ZipFile(archive) as bundle:
names = bundle.namelist()
assert "diagnostics/batch_manifest.json" in names
assert any(name.startswith("outputs/input_00/") for name in names)
assert any(name.startswith("outputs/input_01/") for name in names)
def test_batch_partial_failure_keeps_successful_item_and_records_error(monkeypatch, tmp_path):
job = _job(tmp_path)
prepared = []
for index in range(2):
path = job.work_dir / f"{index:02d}_mix.wav"
path.write_bytes(f"input-{index}".encode())
prepared.append(path)
config = _config(job, prepared)
class FakeSeparator:
def load_model(self, value):
pass
def separate(self, values):
source = Path(values[0])
if source.name.startswith("01_"):
raise RuntimeError("forced second-item failure")
output = job.output_dir / f"{source.stem}_(Vocals).flac"
output.write_bytes(b"success")
return [str(output)]
_patch_runtime(monkeypatch, job, config, FakeSeparator())
result = service.run_prepared_job("state")
status, _, _, outputs, archive, _, _, repro_path = result
assert status.startswith("### Completed with item failures")
assert len(outputs) == 2
stem_outputs = [Path(value) for value in outputs if Path(value).parent == job.output_dir]
assert len(stem_outputs) == 1
assert Path(outputs[-1]).name == "batch_manifest.json"
repro = json.loads(Path(repro_path).read_text(encoding="utf-8"))
assert repro["status"] == "COMPLETED_WITH_ITEM_FAILURES"
batch = repro["execution"]["batch"]
assert batch["completed_count"] == 1
assert batch["failed_count"] == 1
assert batch["items"][1]["status"] == "FAILED"
assert "forced second-item failure" in batch["items"][1]["error"]["message"]
with zipfile.ZipFile(archive) as bundle:
assert any("batch_item_01_error.txt" in name for name in bundle.namelist())
def test_retry_runs_only_failed_item_and_keeps_prior_outputs(monkeypatch, tmp_path):
job = _job(tmp_path)
prepared = []
for index in range(3):
path = job.work_dir / f"{index:02d}_mix.wav"
path.write_bytes(f"input-{index}".encode())
prepared.append(path)
config = _config(job, prepared)
config["inputs"] = [
{
"display_name": f"source-{index}.wav",
"prepared_path": path.relative_to(job.root).as_posix(),
"prepared_sha256": f"sha-{index}",
"prepared_duration_seconds": 10.0,
}
for index, path in enumerate(prepared)
]
config["batch"]["input_count"] = 3
class FakeSeparator:
def __init__(self):
self.load_calls = 0
self.separate_calls = []
def load_model(self, value):
self.load_calls += 1
def separate(self, values):
source = Path(values[0])
self.separate_calls.append(source.name)
if getattr(self, "fail_item_one", False) and source.name.startswith("01_"):
raise RuntimeError("forced retryable item failure")
output = job.output_dir / f"{source.stem}_(Vocals).flac"
output.write_bytes(source.name.encode())
return [str(output)]
first_separator = FakeSeparator()
first_separator.fail_item_one = True
_patch_runtime(monkeypatch, job, config, first_separator)
first = service.run_prepared_job("state")
first_manifest = json.loads((job.config_dir / "batch_manifest.json").read_text(encoding="utf-8"))
assert first[0].startswith("### Completed with item failures")
assert first_manifest["retryable_input_indexes"] == [1]
assert first_manifest["completed_count"] == 2
assert first_separator.separate_calls == [path.name for path in prepared]
second_separator = FakeSeparator()
monkeypatch.setattr(service, "create_separator", lambda **kwargs: second_separator)
second = service.run_prepared_job(
"state", execution_mode=service.EXECUTION_MODE_RETRY_INCOMPLETE
)
final_manifest = json.loads((job.config_dir / "batch_manifest.json").read_text(encoding="utf-8"))
assert second[0].startswith("### Completed")
assert second_separator.separate_calls == [prepared[1].name]
assert final_manifest["status"] == "COMPLETED"
assert final_manifest["completed_count"] == 3
assert final_manifest["retryable_input_indexes"] == []
assert len(final_manifest["attempts"]) == 2
assert final_manifest["attempts"][1]["target_input_indexes"] == [1]
assert final_manifest["total_model_load_count"] == 2
assert sorted(item["input_index"] for item in final_manifest["outputs"]) == [0, 1, 2]
def test_cooperative_stop_marks_remaining_items_retryable(monkeypatch, tmp_path):
job = _job(tmp_path)
prepared = []
for index in range(3):
path = job.work_dir / f"{index:02d}_mix.wav"
path.write_bytes(f"input-{index}".encode())
prepared.append(path)
config = _config(job, prepared)
config["inputs"] = [
{
"display_name": f"source-{index}.wav",
"prepared_path": path.relative_to(job.root).as_posix(),
"prepared_sha256": f"sha-{index}",
"prepared_duration_seconds": 10.0,
}
for index, path in enumerate(prepared)
]
config["batch"]["input_count"] = 3
class FakeSeparator:
def load_model(self, value):
pass
def __init__(self):
self.calls = 0
def separate(self, values):
source = Path(values[0])
output = job.output_dir / f"{source.stem}_(Vocals).flac"
output.write_bytes(source.name.encode())
self.calls += 1
if self.calls == 1:
service._cancel_request_path(job).write_text("{}\n", encoding="utf-8")
return [str(output)]
_patch_runtime(monkeypatch, job, config, FakeSeparator())
result = service.run_prepared_job("state")
manifest = json.loads((job.config_dir / "batch_manifest.json").read_text(encoding="utf-8"))
assert result[0].startswith("### Stopped with partial results")
assert manifest["status"] == "CANCELED_WITH_PARTIAL_RESULTS"
assert manifest["completed_count"] == 1
assert manifest["canceled_count"] == 2
assert manifest["retryable_input_indexes"] == [1, 2]
assert manifest["attempts"][0]["attempted_input_indexes"] == [0]
assert manifest["attempts"][0]["cancel_observed"] is True