Download scripts/run_flow_encoder_pilot.py from coolpoodle/music3lab: direct link, hf CLI and curl.
- Browser
- Download file 6.18 kB
-
https://huggingface.co/coolpoodle/music3lab/resolve/main/scripts/run_flow_encoder_pilot.py
- Command line
-
hf download hf://coolpoodle/music3lab/scripts/run_flow_encoder_pilot.py
-
curl -L -o run_flow_encoder_pilot.py https://huggingface.co/coolpoodle/music3lab/resolve/main/scripts/run_flow_encoder_pilot.py
6.18 kB
| #!/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()) | |