File size: 5,320 Bytes
90884df | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | #!/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())
|