music3lab / scripts /run_native_token_adapter.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
2.85 kB
#!/usr/bin/env python3
"""Run the captured-Music3 native-row adapter pilot."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Sequence
from music3lab.codec.native_token_runner import (
run_native_token_pilot,
verify_native_token_pilot,
verify_token_teacher_dataset,
)
def _parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=(
"Predict native rows only for captured Music3 generations; "
"this is not a native tokenizer or arbitrary-music encoder."
)
)
commands = parser.add_subparsers(dest="command", required=True)
run = commands.add_parser("run")
for name in (
"config",
"flow-config",
"flow-checkpoint",
"snapshot",
"base-manifest",
"diffusers-root",
"dataset-root",
"output-root",
):
run.add_argument(f"--{name}", type=Path, required=True)
verify_data = commands.add_parser("verify-teachers")
verify_data.add_argument("--dataset-root", type=Path, required=True)
verify = commands.add_parser("verify")
verify.add_argument("--output-root", type=Path, required=True)
return parser
def main(argv: Sequence[str] | None = None) -> int:
args = _parser().parse_args(argv)
if args.command == "run":
result = run_native_token_pilot(
config_path=args.config,
flow_config_path=args.flow_config,
flow_checkpoint=args.flow_checkpoint,
snapshot=args.snapshot,
base_manifest=args.base_manifest,
diffusers_root=args.diffusers_root,
dataset_root=args.dataset_root,
output_root=args.output_root,
)
payload = {
"capability": result["metrics"]["capability"],
"dataset_semantic_digest": result["metrics"]["dataset_semantic_digest"],
"metrics_semantic_digest": result["metrics"]["semantic_digest"],
"output_root": str(args.output_root.absolute()),
"native_tokenizer": False,
}
elif args.command == "verify-teachers":
result = verify_token_teacher_dataset(args.dataset_root)
payload = {
"dataset_semantic_digest": result["manifest"]["semantic_digest"],
"split_counts": result["report"].split_counts,
"source_domain": result["report"].source_domain,
}
else:
result = verify_native_token_pilot(args.output_root)
payload = {
"capability": result["metrics"]["capability"],
"metrics_semantic_digest": result["metrics"]["semantic_digest"],
"native_tokenizer": False,
}
print(json.dumps(payload, sort_keys=True))
return 0
if __name__ == "__main__":
raise SystemExit(main())