| |
| """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()) |
|
|