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