music3lab / scripts /run_flow_prepend.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
2.26 kB
#!/usr/bin/env python3
"""Train the bounded local captured-condition prefix adapter."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
def parser() -> argparse.ArgumentParser:
result = argparse.ArgumentParser(
description=(
"Local one-second captured-condition renderer-latent prefix "
"reconstruction only; not song-intro/native-AR prepend, arbitrary "
"WAV, native RVQ, FIM, or generalization."
)
)
for name in (
"config",
"dataset-root",
"snapshot",
"base-manifest",
"diffusers-root",
"output-root",
):
result.add_argument(f"--{name}", required=True)
return result
def main(argv: list[str] | None = None) -> int:
args = parser().parse_args(argv)
from music3lab.editing.flow_prepend_pilot import train_flow_prepend_pilot
metrics = train_flow_prepend_pilot(
config_path=Path(args.config),
dataset_root=Path(args.dataset_root),
snapshot=Path(args.snapshot),
base_manifest=Path(args.base_manifest),
diffusers_root=Path(args.diffusers_root),
output_root=Path(args.output_root),
)
print(
json.dumps(
{
"measured_improvement_gate": metrics[
"measured_improvement_gate"
],
"median_prefix_latent_nmse_improvement_fraction": metrics[
"median_prefix_latent_nmse_improvement_fraction"
],
"median_prefix_audio_ruler_improvement_fraction": metrics[
"median_prefix_audio_ruler_improvement_fraction"
],
"right_suffix_latent_exact": metrics[
"right_suffix_latent_exact"
],
"composited_suffix_audio_exact": metrics[
"composited_suffix_audio_exact"
],
"output_root": metrics["output_root"],
"semantic_digest": metrics["semantic_digest"],
},
sort_keys=True,
)
)
return 0 if metrics["measured_improvement_gate"] else 2
if __name__ == "__main__":
raise SystemExit(main())