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