music3lab / scripts /run_flow_encoder_pilot.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
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())