music3lab / scripts /run_learned_audio_continuation.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
3.08 kB
#!/usr/bin/env python3
"""Run the bounded learned continuous-latent continuation pilot."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Sequence
from music3lab.editing.learned_audio_continuation_data import (
load_learned_continuation_config,
load_pair_records,
)
from music3lab.editing.learned_audio_continuation_runner import (
run_learned_audio_continuation,
)
def _parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser()
commands = parser.add_subparsers(dest="command", required=True)
check = commands.add_parser("check")
check.add_argument("--config", type=Path, required=True)
check.add_argument("--corpus-root", type=Path, required=True)
run = commands.add_parser("run")
run.add_argument("--config", type=Path, required=True)
run.add_argument("--flow-config", type=Path, required=True)
run.add_argument("--corpus-root", type=Path, required=True)
run.add_argument("--encoder-checkpoint", type=Path, required=True)
run.add_argument("--audio", type=Path, required=True)
run.add_argument("--snapshot", type=Path, required=True)
run.add_argument("--base-manifest", type=Path, required=True)
run.add_argument("--diffusers-root", type=Path, required=True)
run.add_argument("--output-root", type=Path, required=True)
return parser
def main(argv: Sequence[str] | None = None) -> int:
arguments = _parser().parse_args(argv)
if arguments.command == "check":
loaded = load_learned_continuation_config(arguments.config)
records = load_pair_records(arguments.corpus_root, loaded)
result = {
"schema_version": loaded.config.schema_version,
"config_file_sha256": loaded.file_sha256,
"split_counts": {key: len(value) for key, value in records.items()},
"target_visible_to_conditioner": False,
"specialist_champion_eligible": False,
}
status = 0
elif arguments.command == "run":
metrics = run_learned_audio_continuation(
config_path=arguments.config,
flow_config_path=arguments.flow_config,
corpus_root=arguments.corpus_root,
encoder_checkpoint=arguments.encoder_checkpoint,
demonstration_audio=arguments.audio,
snapshot=arguments.snapshot,
base_manifest=arguments.base_manifest,
diffusers_root=arguments.diffusers_root,
output_root=arguments.output_root,
)
result = {
"output_root": metrics["output_root"],
"status": metrics["status"],
"best_step": metrics["best_step"],
"heldout_gate": metrics["heldout"]["passes"],
"specialist_promoted": False,
}
status = 0 if metrics["status"] == "FUNCTIONAL_PASS" else 1
else:
raise AssertionError("unreachable")
print(json.dumps(result, sort_keys=True, separators=(",", ":")))
return status
if __name__ == "__main__":
raise SystemExit(main())