music3lab / native_token_compatibility.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
8.89 kB
"""Validate and exactly compare native MiniMax Music 3 token trajectories."""
from __future__ import annotations
import argparse
import json
import pickle
from pathlib import Path
from typing import Any, Literal
import numpy as np
import torch
TokenLayout = Literal["frames_first", "codebooks_first"]
LAYOUTS: tuple[TokenLayout, ...] = ("frames_first", "codebooks_first")
NUM_CODEBOOKS = 8
CODEBOOK_SIZES = (16_384, 1_024, 1_024, 1_024, 1_024, 1_024, 1_024, 1_024)
FRAME_RATE_HZ = 25
def normalize_tokens(tokens: Any, *, layout: TokenLayout) -> np.ndarray:
"""Validate layout explicitly and return a CPU NumPy array ``[frames, 8]``."""
if layout not in LAYOUTS:
raise ValueError(f"layout must be one of {LAYOUTS}, got {layout!r}")
if isinstance(tokens, torch.Tensor):
if tokens.device.type != "cpu":
raise ValueError("token comparison is CPU-only; pass a CPU tensor")
array = tokens.detach().numpy()
else:
array = np.asarray(tokens)
if array.ndim != 2:
raise ValueError(f"tokens must be rank 2, got shape {array.shape}")
expected_axis = 1 if layout == "frames_first" else 0
if array.shape[expected_axis] != NUM_CODEBOOKS:
expected = "[frames, 8]" if layout == "frames_first" else "[8, frames]"
raise ValueError(f"{layout} tokens must have shape {expected}, got {array.shape}")
if not np.issubdtype(array.dtype, np.integer) or np.issubdtype(array.dtype, np.bool_):
raise TypeError(f"tokens must use an integer dtype, got {array.dtype}")
frames_first = array if layout == "frames_first" else array.T
if frames_first.shape[0] == 0:
raise ValueError("token trajectory must contain at least one frame")
return np.ascontiguousarray(frames_first, dtype=np.int64)
def validate_native_tokens(tokens: Any, *, layout: TokenLayout) -> np.ndarray:
"""Return normalized tokens after enforcing every native codebook range."""
normalized = normalize_tokens(tokens, layout=layout)
for index, vocab_size in enumerate(CODEBOOK_SIZES):
values = normalized[:, index]
minimum = int(values.min())
maximum = int(values.max())
if minimum < 0 or maximum >= vocab_size:
raise ValueError(
f"codebook c{index} values must be in [0, {vocab_size - 1}], "
f"observed min={minimum}, max={maximum}"
)
return normalized
def compare_native_tokens(
reference: Any,
recovered: Any,
*,
reference_layout: TokenLayout,
recovered_layout: TokenLayout,
) -> dict[str, Any]:
"""Compare trajectories without truncation, alignment, or tolerance."""
reference_tokens = validate_native_tokens(reference, layout=reference_layout)
recovered_tokens = validate_native_tokens(recovered, layout=recovered_layout)
shape_match = reference_tokens.shape == recovered_tokens.shape
frames = int(reference_tokens.shape[0]) if shape_match else None
per_codebook: list[dict[str, Any]] = []
for index, vocab_size in enumerate(CODEBOOK_SIZES):
entry: dict[str, Any] = {"codebook": index, "vocab_size": vocab_size}
if shape_match:
matches = int(np.count_nonzero(reference_tokens[:, index] == recovered_tokens[:, index]))
entry.update(
{
"matching_frames": matches,
"total_frames": frames,
"agreement": matches / frames,
"agreement_percent": 100.0 * matches / frames,
"exact": matches == frames,
}
)
else:
entry.update(
{
"matching_frames": None,
"total_frames": None,
"agreement": None,
"agreement_percent": None,
"exact": False,
}
)
per_codebook.append(entry)
exact_match = shape_match and bool(np.array_equal(reference_tokens, recovered_tokens))
return {
"status": "MATCH" if exact_match else "MISMATCH",
"exact_match": exact_match,
"shape_match": shape_match,
"reference_shape_frames_first": list(reference_tokens.shape),
"recovered_shape_frames_first": list(recovered_tokens.shape),
"frames": frames,
"duration_seconds_at_25hz": frames / FRAME_RATE_HZ if frames is not None else None,
"frame_rate_hz": FRAME_RATE_HZ,
"codebook_sizes": list(CODEBOOK_SIZES),
"per_codebook": per_codebook,
}
def _select_loaded(value: Any, *, key: str | None, source: Path) -> Any:
if isinstance(value, np.lib.npyio.NpzFile):
available = list(value.files)
selected = key or (available[0] if len(available) == 1 else None)
if selected is None or selected not in available:
raise ValueError(f"{source} contains arrays {available}; select one with --*-key")
return value[selected]
if isinstance(value, dict):
available = list(value)
selected = key or (available[0] if len(available) == 1 else None)
if selected is None or selected not in value:
raise ValueError(f"{source} contains keys {available}; select one with --*-key")
return value[selected]
if key is not None:
raise ValueError(f"{source} is a bare array/tensor, so a key cannot be selected")
return value
def load_token_file(path: str | Path, *, key: str | None = None) -> Any:
source = Path(path)
suffix = source.suffix.lower()
if suffix in {".pt", ".pth"}:
try:
value = torch.load(str(source), map_location="cpu", weights_only=True)
except (pickle.UnpicklingError, EOFError, OSError, RuntimeError) as error:
raise ValueError(f"unable to safely load token file {source}: {error}") from error
return _select_loaded(value, key=key, source=source)
if suffix == ".json":
with source.open("r", encoding="utf-8") as handle:
value = json.load(handle)
return _select_loaded(value, key=key, source=source)
if suffix in {".npy", ".npz"}:
value = np.load(source, allow_pickle=False)
try:
selected = _select_loaded(value, key=key, source=source)
return np.asarray(selected).copy() if isinstance(value, np.lib.npyio.NpzFile) else selected
finally:
if isinstance(value, np.lib.npyio.NpzFile):
value.close()
raise ValueError(f"unsupported token file extension {suffix!r}; use .npy, .npz, .pt, .pth, or .json")
def format_human(result: dict[str, Any]) -> str:
lines = [
f"Status: {result['status']}",
f"Reference shape [frames, codebooks]: {result['reference_shape_frames_first']}",
f"Recovered shape [frames, codebooks]: {result['recovered_shape_frames_first']}",
f"Frame rate: {result['frame_rate_hz']} Hz",
]
if result["duration_seconds_at_25hz"] is not None:
lines.append(f"Duration represented: {result['duration_seconds_at_25hz']:.3f} s")
for entry in result["per_codebook"]:
agreement = entry["agreement_percent"]
rendered = "not comparable (shape mismatch)" if agreement is None else f"{agreement:.6f}%"
lines.append(
f"CB{entry['codebook']} [0, {entry['vocab_size'] - 1}]: {rendered}"
)
return "\n".join(lines)
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("reference", help="Internal Music 3 token trajectory")
parser.add_argument("recovered", help="Tokens recovered by the candidate encoder")
parser.add_argument("--reference-layout", required=True, choices=LAYOUTS)
parser.add_argument("--recovered-layout", required=True, choices=LAYOUTS)
parser.add_argument("--reference-key")
parser.add_argument("--recovered-key")
parser.add_argument("--json", action="store_true")
return parser
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
reference = load_token_file(args.reference, key=args.reference_key)
recovered = load_token_file(args.recovered, key=args.recovered_key)
result = compare_native_tokens(
reference,
recovered,
reference_layout=args.reference_layout,
recovered_layout=args.recovered_layout,
)
except (FileNotFoundError, OSError, RuntimeError, TypeError, ValueError) as error:
payload = {"status": "ERROR", "error": str(error)}
print(json.dumps(payload, indent=2) if args.json else f"ERROR: {error}")
return 2
print(json.dumps(result, indent=2) if args.json else format_human(result))
return 0 if result["exact_match"] else 1
if __name__ == "__main__":
raise SystemExit(main())