music3lab / encode_audio.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
5.6 kB
"""Fail-fast native-token encoding API for the released Music 3 checkpoint.
The public checkpoint does not contain the RVQ quantizer/codebooks required by
this operation. This module exists so callers get a precise, reusable failure
instead of accidentally treating continuous DAV latents as discrete tokens.
"""
from __future__ import annotations
import argparse
import json
import os
from pathlib import Path
from typing import Any
import torch
from inspect_dav import inspect_checkpoint
DAV_PATH_ENV = "MINIMAX_DAV_PATH"
class NativeTokenizerUnavailableError(RuntimeError):
"""The supplied release checkpoint cannot produce native Music 3 tokens."""
def __init__(self, message: str, *, report: dict[str, Any] | None = None) -> None:
super().__init__(message)
self.report = report
def resolve_dav_path(dav_path: str | os.PathLike[str] | None = None) -> Path:
value = dav_path if dav_path is not None else os.environ.get(DAV_PATH_ENV)
if value is None or not str(value).strip():
raise NativeTokenizerUnavailableError(
f"no DAV checkpoint supplied; pass dav_path or set {DAV_PATH_ENV}"
)
return Path(value).expanduser()
def require_native_tokenizer(
dav_path: str | os.PathLike[str] | None = None,
) -> dict[str, Any]:
"""Inspect ``dav.pth`` and return its report only if a native tokenizer exists."""
resolved = resolve_dav_path(dav_path)
report = inspect_checkpoint(resolved)
capabilities = report["capabilities"]
if not capabilities["can_encode_native_music3_tokens"]:
missing = [
label
for label, available in (
("waveform/tokenizer encoder weights", capabilities["waveform_analysis_encoder_weights"]),
("RVQ/VQ quantizer weights", capabilities["rvq_or_vq_quantizer_weights"]),
("codebook embedding weights", capabilities["codebook_embedding_weights"]),
("serialized tokenizer architecture/config", capabilities["serialized_architecture_config"]),
)
if not available
]
if missing:
reason = "released dav.pth cannot encode native Music 3 tokens: missing " + ", ".join(missing) + "."
else:
reason = (
"released dav.pth has tokenizer-like candidate weights/config, but no exact compatible "
"executable Music 3 tokenizer architecture/API has been implemented and verified."
)
continuous_evidence = (
capabilities["waveform_analysis_encoder_weights"]
and capabilities["continuous_gaussian_posterior_heads"]
and capabilities["continuous_flow_weights"]
)
if continuous_evidence:
reason += (
" encoder.* plus mean_proj.* and logs_proj.* form a continuous "
"Flow-VAE analysis path; they do not emit integer RVQ codes."
)
raise NativeTokenizerUnavailableError(reason, report=report)
return report
def encode_audio(
audio_path: str | os.PathLike[str],
*,
dav_path: str | os.PathLike[str] | None = None,
) -> torch.Tensor:
"""Return native tokens as ``[frames, 8]`` or fail before reading the WAV.
The release inspected for this project always takes the failure path. The
return annotation records the intended contract without manufacturing token
IDs or using a third-party codec.
"""
# Capability inspection intentionally precedes even checking the input path.
require_native_tokenizer(dav_path)
raise NativeTokenizerUnavailableError(
"checkpoint advertises tokenizer-like weights, but no released Music 3 "
"tokenizer architecture/API is available to execute them safely"
)
def _blocked_payload(error: NativeTokenizerUnavailableError, audio_path: str) -> dict[str, Any]:
payload: dict[str, Any] = {
"status": "BLOCKED",
"operation": "encode_audio",
"audio_path": audio_path,
"audio_was_read": False,
"reason": str(error),
}
if error.report is not None:
payload["capabilities"] = error.report["capabilities"]
payload["state_dict"] = error.report["state_dict"]
return payload
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("audio", help="Input WAV path (not opened when capability is blocked)")
parser.add_argument("--dav", help=f"Path to dav.pth; defaults to {DAV_PATH_ENV}")
parser.add_argument("--json", action="store_true", help="Emit machine-readable JSON")
return parser
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
tokens = encode_audio(args.audio, dav_path=args.dav)
except (NativeTokenizerUnavailableError, FileNotFoundError, RuntimeError, ValueError) as error:
if isinstance(error, NativeTokenizerUnavailableError):
native_error = error
else:
native_error = NativeTokenizerUnavailableError(str(error))
payload = _blocked_payload(native_error, args.audio)
print(json.dumps(payload, indent=2) if args.json else f"BLOCKED: {payload['reason']}")
return 2
# Kept for API completeness if an official tokenizer is released later.
payload = {"status": "OK", "shape": list(tokens.shape), "dtype": str(tokens.dtype)}
print(json.dumps(payload, indent=2) if args.json else f"tokens: {tuple(tokens.shape)}")
return 0
if __name__ == "__main__":
raise SystemExit(main())