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