#!/usr/bin/env python3 """Run the bounded continuous flow-latent encoder pilot.""" from __future__ import annotations import argparse from pathlib import Path from typing import Sequence from music3lab.codec.flow_encoder import canonical_json_bytes from music3lab.codec.runner import ( generate_teacher_dataset, load_teacher_dataset, recover_teacher_dataset, train_flow_encoder, verify_pilot_bundle, ) def _common(parser: argparse.ArgumentParser) -> None: parser.add_argument("--config", type=Path, required=True) parser.add_argument("--snapshot", type=Path, required=True) parser.add_argument("--base-manifest", type=Path, required=True) parser.add_argument("--diffusers-root", type=Path, required=True) def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description=( "Predict continuous Music 3 Flow-VAE renderer latents from " "waveforms; this does not produce native RVQ tokens." ) ) commands = parser.add_subparsers(dest="command", required=True) generate = commands.add_parser("generate-teachers") _common(generate) generate.add_argument("--dataset-root", type=Path, required=True) recover = commands.add_parser("recover-teachers") _common(recover) recover.add_argument("--quarantine-root", type=Path, required=True) recover.add_argument("--producer-commit", required=True) recover.add_argument("--dataset-root", type=Path, required=True) train = commands.add_parser("train") _common(train) train.add_argument("--dataset-root", type=Path, required=True) train.add_argument("--external-wav", type=Path, required=True) train.add_argument("--output-root", type=Path, required=True) run = commands.add_parser("run") _common(run) run.add_argument("--dataset-root", type=Path, required=True) run.add_argument("--external-wav", type=Path, required=True) run.add_argument("--output-root", type=Path, required=True) verify_data = commands.add_parser("verify-teachers") verify_data.add_argument("--dataset-root", type=Path, required=True) verify_run = commands.add_parser("verify") verify_run.add_argument("--output-root", type=Path, required=True) return parser def _generate(arguments: argparse.Namespace) -> dict[str, object]: manifest = generate_teacher_dataset( config_path=arguments.config, snapshot=arguments.snapshot, base_manifest=arguments.base_manifest, diffusers_root=arguments.diffusers_root, output_root=arguments.dataset_root, ) return { "kind": "continuous_flow_latent_teachers", "native_rvq": False, "dataset_root": str(arguments.dataset_root.absolute()), "manifest_semantic_digest": manifest.semantic_digest, "split_counts": { key: value.count for key, value in manifest.splits.items() }, } def _recover(arguments: argparse.Namespace) -> dict[str, object]: manifest = recover_teacher_dataset( config_path=arguments.config, quarantine_root=arguments.quarantine_root, producer_project_git_commit=arguments.producer_commit, snapshot=arguments.snapshot, base_manifest=arguments.base_manifest, diffusers_root=arguments.diffusers_root, output_root=arguments.dataset_root, ) return { "kind": "continuous_flow_latent_teachers", "native_rvq": False, "dataset_root": str(arguments.dataset_root.absolute()), "manifest_semantic_digest": manifest.semantic_digest, "producer_project_git_commit": manifest.producer_project_git_commit, "publication_project_git_commit": ( manifest.publication_project_git_commit ), "recovered_from_complete_quarantine": True, "replay_exact_count": manifest.replay_exact_count, } def _train(arguments: argparse.Namespace) -> dict[str, object]: metrics = train_flow_encoder( config_path=arguments.config, dataset_root=arguments.dataset_root, snapshot=arguments.snapshot, base_manifest=arguments.base_manifest, diffusers_root=arguments.diffusers_root, external_wav=arguments.external_wav, output_root=arguments.output_root, ) return { "kind": "continuous_flow_latent_encoder_pilot", "native_rvq_capability": metrics.native_rvq_capability, "measured_improvement_gate": metrics.measured_improvement_gate, "metrics_semantic_digest": metrics.semantic_digest, "output_root": str(arguments.output_root.absolute()), } def main(argv: Sequence[str] | None = None) -> int: arguments = _parser().parse_args(argv) if arguments.command == "generate-teachers": result = _generate(arguments) elif arguments.command == "recover-teachers": result = _recover(arguments) elif arguments.command == "train": result = _train(arguments) elif arguments.command == "run": dataset = _generate(arguments) result = {"dataset": dataset, "pilot": _train(arguments)} elif arguments.command == "verify-teachers": dataset = load_teacher_dataset(arguments.dataset_root) result = { "kind": "continuous_flow_latent_teachers", "native_rvq": False, "manifest_semantic_digest": dataset.manifest.semantic_digest, "split_counts": { key: value.count for key, value in dataset.manifest.splits.items() }, } elif arguments.command == "verify": metrics = verify_pilot_bundle(arguments.output_root) result = { "kind": "continuous_flow_latent_encoder_pilot", "native_rvq_capability": metrics.native_rvq_capability, "measured_improvement_gate": metrics.measured_improvement_gate, "metrics_semantic_digest": metrics.semantic_digest, } else: raise AssertionError("unreachable command") print(canonical_json_bytes(result).decode("utf-8"), end="") return 0 if __name__ == "__main__": raise SystemExit(main())