#!/usr/bin/env python3 """Freeze external splits and run the bounded continuous-encoder fine-tune.""" from __future__ import annotations import argparse import json from pathlib import Path from typing import Sequence from music3lab.codec.external_finetune import ( InterimExternalDataConfig, load_external_finetune_config, load_interim_source_splits, load_source_splits, ) from music3lab.codec.external_finetune_runner import ( freeze_external_source_splits, run_external_finetune, ) def _parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description=( "Fine-tune the continuous Flow renderer-latent encoder on external " "audio; this does not recover native Music3 RVQ tokens." ) ) commands = parser.add_subparsers(dest="command", required=True) config = commands.add_parser("check-config") config.add_argument("--config", type=Path, required=True) freeze = commands.add_parser("freeze-splits") freeze.add_argument("--config", type=Path, required=True) freeze.add_argument("--accepted-manifest", type=Path, required=True) freeze.add_argument("--user22-hashes", type=Path, required=True) freeze.add_argument("--output", type=Path, required=True) verify = commands.add_parser("check-splits") verify.add_argument("--splits", type=Path, required=True) verify.add_argument("--config", type=Path) train = commands.add_parser("train") train.add_argument("--config", type=Path, required=True) train.add_argument("--flow-config", type=Path, required=True) train.add_argument("--initial-checkpoint", type=Path, required=True) train.add_argument("--external-splits", type=Path, required=True) train.add_argument("--user22-hashes", type=Path, required=True) train.add_argument("--user22-canonical", type=Path, required=True) train.add_argument("--teacher-dataset", type=Path, required=True) train.add_argument("--snapshot", type=Path, required=True) train.add_argument("--base-manifest", type=Path, required=True) train.add_argument("--diffusers-root", type=Path, required=True) train.add_argument("--output-root", type=Path, required=True) return parser def main(argv: Sequence[str] | None = None) -> int: arguments = _parser().parse_args(argv) if arguments.command == "check-config": loaded = load_external_finetune_config(arguments.config) result = { "schema_version": loaded.config.schema_version, "config_file_sha256": loaded.file_sha256, "config_semantic_digest": loaded.semantic_digest, "steps": loaded.config.training.steps, "batch_composition": { "external": loaded.config.training.external_batch_size, "music3": loaded.config.training.music3_batch_size, }, "native_tokenizer": False, } elif arguments.command == "freeze-splits": splits = freeze_external_source_splits( accepted_manifest=arguments.accepted_manifest, user22_hashes_manifest=arguments.user22_hashes, config_path=arguments.config, output_path=arguments.output, ) result = { "split_manifest": str(arguments.output.absolute()), "counts": {name: len(rows) for name, rows in splits.items()}, "user22_in_development_splits": False, } elif arguments.command == "check-splits": if arguments.config is None: splits = load_source_splits(arguments.splits) else: loaded = load_external_finetune_config(arguments.config) if not isinstance(loaded.config.data, InterimExternalDataConfig): raise ValueError("--config is only required for interim splits") splits = load_interim_source_splits( arguments.splits, loaded.config.data ) result = { "split_manifest": str(arguments.splits.resolve(strict=True)), "counts": {name: len(rows) for name, rows in splits.items()}, } elif arguments.command == "train": metrics = run_external_finetune( config_path=arguments.config, flow_config_path=arguments.flow_config, initial_checkpoint=arguments.initial_checkpoint, external_splits_path=arguments.external_splits, user22_hashes_path=arguments.user22_hashes, user22_canonical_path=arguments.user22_canonical, teacher_dataset_root=arguments.teacher_dataset, snapshot=arguments.snapshot, base_manifest=arguments.base_manifest, diffusers_root=arguments.diffusers_root, output_root=arguments.output_root, ) result = { "output_root": str(arguments.output_root.absolute()), "selected_checkpoint": metrics["selected_checkpoint"], "heldout_gate": metrics["heldout_gate"]["passes"], "native_tokenizer": False, "generalization_claim": False, } else: raise AssertionError("unreachable command") print(json.dumps(result, sort_keys=True, separators=(",", ":"))) return 0 if __name__ == "__main__": raise SystemExit(main())