File size: 3,079 Bytes
90884df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
#!/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())