diff --git a/examples/moss_sound_effects.py b/examples/moss_sound_effects.py new file mode 100644 index 000000000..e2b2ff193 --- /dev/null +++ b/examples/moss_sound_effects.py @@ -0,0 +1,97 @@ +#!/usr/bin/env python +"""Sound-effect generation with MOSS-SoundEffect. + +This example demonstrates ambient and event-conditioned generation for the +SoundEffect variant. + +Usage: + uv run python examples/moss_sound_effects.py + uv run python examples/moss_sound_effects.py --ambient-sound "busy cafe" --sound-event "cups clinking" +""" + +from __future__ import annotations + +import argparse + +from mlx_audio.tts.generate import generate_audio + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="MOSS-SoundEffect example") + parser.add_argument( + "--model", + default="OpenMOSS-Team/MOSS-SoundEffect", + help="SoundEffect model id or local path", + ) + parser.add_argument( + "--ambient-sound", + default="Thunder and rain over a city street.", + help="Ambient scene description", + ) + parser.add_argument( + "--sound-event", + default="storm", + help="Primary sound event cue", + ) + parser.add_argument( + "--quality", + default="high", + help="Quality hint", + ) + parser.add_argument( + "--preset", + default="soundeffect", + help="Sampling preset", + ) + parser.add_argument( + "--tokens", + type=int, + default=140, + help="Target token budget", + ) + parser.add_argument( + "--max-tokens", + type=int, + default=260, + help="Safety cap for generation", + ) + parser.add_argument( + "--output-dir", + default="outputs/moss_sound_effects", + help="Directory for generated outputs", + ) + parser.add_argument( + "--file-prefix", + default="sound_effect", + help="Output filename prefix", + ) + parser.add_argument( + "--verbose", + action="store_true", + help="Print runtime stats", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + # `text` is intentionally omitted: generate_audio() mirrors CLI behavior and + # backfills `text` from `ambient_sound` when needed for SoundEffect prompts. + generate_audio( + text=None, + model=args.model, + preset=args.preset, + ambient_sound=args.ambient_sound, + sound_event=args.sound_event, + quality=args.quality, + tokens=args.tokens, + max_tokens=args.max_tokens, + output_path=args.output_dir, + file_prefix=args.file_prefix, + verbose=args.verbose, + play=False, + ) + + +if __name__ == "__main__": + main() diff --git a/examples/moss_ttsd_dialogue.py b/examples/moss_ttsd_dialogue.py new file mode 100644 index 000000000..61084af80 --- /dev/null +++ b/examples/moss_ttsd_dialogue.py @@ -0,0 +1,163 @@ +#!/usr/bin/env python +"""Multi-speaker dialogue generation with MOSS-TTSD. + +This example demonstrates `dialogue_speakers` schema usage and speaker-tagged +prompting (`[S1]`, `[S2]`, ...). + +Usage: + uv run python examples/moss_ttsd_dialogue.py + uv run python examples/moss_ttsd_dialogue.py --dialogue-speakers-json /path/to/speakers.json + uv run python examples/moss_ttsd_dialogue.py --use-zero-based-speaker-ids +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +from mlx_audio.tts.generate import generate_audio + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="MOSS-TTSD multi-speaker dialogue example" + ) + parser.add_argument( + "--model", + default="OpenMOSS-Team/MOSS-TTSD-v1.0", + help="TTSD model id or local path", + ) + parser.add_argument( + "--text", + default=( + "[S1] Thanks for joining this dialogue demo. " + "[S2] Happy to help, we are validating speaker switching and continuity." + ), + help="Dialogue text with [S#] speaker tags", + ) + parser.add_argument( + "--dialogue-speakers-json", + default=None, + help="Optional path to speaker schema JSON list", + ) + parser.add_argument( + "--speaker1-ref-audio", + default="REFERENCE/MOSS-Audio-Tokenizer/demo/demo_gt.wav", + help="Fallback speaker 1 reference audio", + ) + parser.add_argument( + "--speaker1-ref-text", + default="Speaker one reference prompt.", + help="Fallback speaker 1 transcript", + ) + parser.add_argument( + "--speaker2-ref-audio", + default="REFERENCE/MOSS-Audio-Tokenizer/demo/demo_gt.wav", + help="Fallback speaker 2 reference audio", + ) + parser.add_argument( + "--speaker2-ref-text", + default="Speaker two reference prompt.", + help="Fallback speaker 2 transcript", + ) + parser.add_argument( + "--use-zero-based-speaker-ids", + action="store_true", + help="Emit fallback schema with IDs [0,1] instead of [1,2]", + ) + parser.add_argument( + "--preset", + default="ttsd", + help="Sampling preset", + ) + parser.add_argument( + "--input-type", + choices=["text", "pinyin", "ipa"], + default="text", + help="Input representation", + ) + parser.add_argument( + "--language", + default="en", + help="Optional language hint", + ) + parser.add_argument( + "--tokens", + type=int, + default=200, + help="Target token budget", + ) + parser.add_argument( + "--max-tokens", + type=int, + default=320, + help="Safety cap for generation", + ) + parser.add_argument( + "--output-dir", + default="outputs/moss_ttsd_dialogue", + help="Directory for generated outputs", + ) + parser.add_argument( + "--file-prefix", + default="dialogue", + help="Output filename prefix", + ) + parser.add_argument( + "--verbose", + action="store_true", + help="Print runtime stats", + ) + return parser.parse_args() + + +def build_fallback_schema(args: argparse.Namespace, output_dir: Path) -> Path: + output_dir.mkdir(parents=True, exist_ok=True) + first_id = 0 if args.use_zero_based_speaker_ids else 1 + second_id = first_id + 1 + + schema = [ + { + "speaker_id": first_id, + "ref_audio": args.speaker1_ref_audio, + "ref_text": args.speaker1_ref_text, + }, + { + "speaker_id": second_id, + "ref_audio": args.speaker2_ref_audio, + "ref_text": args.speaker2_ref_text, + }, + ] + + schema_path = output_dir / "dialogue_speakers_demo.json" + schema_path.write_text(json.dumps(schema, indent=2), encoding="utf-8") + return schema_path + + +def main() -> None: + args = parse_args() + output_dir = Path(args.output_dir) + dialogue_speakers_json = args.dialogue_speakers_json + if dialogue_speakers_json is None: + dialogue_speakers_json = str(build_fallback_schema(args, output_dir)) + print(f"Wrote fallback dialogue schema: {dialogue_speakers_json}") + + generate_audio( + text=args.text, + model=args.model, + preset=args.preset, + input_type=args.input_type, + language=args.language, + dialogue_speakers_json=dialogue_speakers_json, + tokens=args.tokens, + max_tokens=args.max_tokens, + output_path=str(output_dir), + file_prefix=args.file_prefix, + verbose=args.verbose, + play=False, + ) + + +if __name__ == "__main__": + main() diff --git a/examples/moss_voice_design.py b/examples/moss_voice_design.py new file mode 100644 index 000000000..23a2c320b --- /dev/null +++ b/examples/moss_voice_design.py @@ -0,0 +1,121 @@ +#!/usr/bin/env python +"""Voice design with MOSS-Voice-Generator. + +This example demonstrates description-first voice creation using `instruct` +(voice design prompt), plus optional language/quality controls. + +Usage: + uv run python examples/moss_voice_design.py + uv run python examples/moss_voice_design.py --instruct "Energetic sports commentator" + uv run python examples/moss_voice_design.py --text "ni hao" --input-type pinyin +""" + +from __future__ import annotations + +import argparse + +from mlx_audio.tts.generate import generate_audio + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="MOSS-Voice-Generator example") + parser.add_argument( + "--model", + default="OpenMOSS-Team/MOSS-Voice-Generator", + help="VoiceGenerator model id or local path", + ) + parser.add_argument( + "--text", + default="Welcome to the voice design example running on MLX.", + help="Text to synthesize", + ) + parser.add_argument( + "--instruct", + default="A calm, clear studio narrator with warm tone.", + help="Voice design instruction prompt", + ) + parser.add_argument( + "--preset", + default="voice_generator", + help="Sampling preset", + ) + parser.add_argument( + "--input-type", + choices=["text", "pinyin", "ipa"], + default="text", + help="Input representation", + ) + parser.add_argument( + "--language", + default="en", + help="Optional language hint", + ) + parser.add_argument( + "--quality", + default="high", + help="Quality hint (draft/balanced/high/max/custom:...)", + ) + parser.add_argument( + "--tokens", + type=int, + default=180, + help="Target token budget", + ) + parser.add_argument( + "--duration-s", + type=float, + default=None, + help="Optional duration hint in seconds", + ) + parser.add_argument( + "--max-tokens", + type=int, + default=300, + help="Safety cap for generation", + ) + parser.add_argument( + "--normalize-inputs", + action="store_true", + help="Force text/instruction normalization before prompt packing", + ) + parser.add_argument( + "--output-dir", + default="outputs/moss_voice_design", + help="Directory for generated outputs", + ) + parser.add_argument( + "--file-prefix", + default="voice_design", + help="Output filename prefix", + ) + parser.add_argument( + "--verbose", + action="store_true", + help="Print runtime stats", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + generate_audio( + text=args.text, + model=args.model, + preset=args.preset, + input_type=args.input_type, + language=args.language, + quality=args.quality, + instruct=args.instruct, + tokens=args.tokens, + duration_s=args.duration_s, + max_tokens=args.max_tokens, + normalize_inputs=args.normalize_inputs, + output_path=args.output_dir, + file_prefix=args.file_prefix, + verbose=args.verbose, + play=False, + ) + + +if __name__ == "__main__": + main() diff --git a/mlx_audio/codec/__init__.py b/mlx_audio/codec/__init__.py index 8b4bdd72b..4088ee6f8 100644 --- a/mlx_audio/codec/__init__.py +++ b/mlx_audio/codec/__init__.py @@ -1 +1 @@ -from .models import DAC, Encodec, Mimi, Vocos +from .models import DAC, Encodec, Mimi, MossAudioTokenizer, Vocos diff --git a/mlx_audio/codec/models/__init__.py b/mlx_audio/codec/models/__init__.py index 3180809c8..9664477b1 100644 --- a/mlx_audio/codec/models/__init__.py +++ b/mlx_audio/codec/models/__init__.py @@ -1,5 +1,6 @@ from .descript import DAC from .encodec import Encodec from .mimi import Mimi +from .moss_audio_tokenizer import MossAudioTokenizer from .snac import SNAC from .vocos import Vocos diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/README.md b/mlx_audio/codec/models/moss_audio_tokenizer/README.md new file mode 100644 index 000000000..efdbfdde9 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/README.md @@ -0,0 +1,88 @@ +# MOSS Audio Tokenizer (Shared Codec) + +Shared codec used by all MOSS-TTS runtimes. + +- Package: `mlx_audio/codec/models/moss_audio_tokenizer/` +- Runtime class: `MossAudioTokenizer` +- Typical frame rate: `sampling_rate / downsample_rate` (24 kHz / 1920 => 12.5 Hz for upstream checkpoints) + +## Why This Matters For MOSS-TTS + +`moss_tts` and `moss_tts_realtime` both depend on this codec for: + +- reference-audio encoding into discrete codebook tokens +- generated token decoding back to waveform +- streaming decode paths with cache boundaries + +If codec contracts are wrong, generation may "run" but output quality/continuity will degrade. + +## Core APIs + +- `encode(input_values, ..., num_quantizers=None, chunk_duration=None)` +- `decode(audio_codes, ..., num_quantizers=None, chunk_duration=None)` +- `batch_encode(wav_list, num_quantizers=None)` +- `batch_decode(codes_list, num_quantizers=None)` +- `streaming_decode(audio_codes, chunk_tokens=..., num_quantizers=None)` + +## Audio-Code Shape Contracts + +`decode(...)` accepts both canonical and transposed layouts: + +- 2D: `(NQ, T)` or `(T, NQ)` +- 3D: `(NQ, B, T)` or `(B, T, NQ)` + +Internally, decode normalizes to canonical `(NQ, B, T)`. + +`batch_decode(...)` expects each list entry to resolve to `batch_size=1` after normalization. + +`streaming_decode(...)` currently supports only `batch_size=1`. + +## `num_quantizers` Behavior + +- If `num_quantizers` is omitted, runtime uses configured quantizer count. +- If provided, it must be `1..configured_nq`. +- Prefix decode is supported (decode with fewer quantizers than checkpoint max). + +Ambiguous shape ties are resolved with conservative rules favoring canonical orientation to preserve encode->decode round-trip behavior. + +## Chunked Encode/Decode + +When `chunk_duration` is used: + +- must be positive +- must be `<= causal_transformer_context_duration` +- `chunk_duration * sampling_rate` must be divisible by `downsample_rate` +- streaming chunked paths currently require `batch_size=1` + +Chunked decode and `streaming_decode(...)` include explicit `mx.eval(...)`/`mx.clear_cache()` boundaries to keep long runs bounded. + +## Checkpoint Sanitization + +`sanitize(...)` performs two key operations: + +1. Tensor layout transpose when checkpoint 3D layouts differ from current MLX parameter shapes. +2. Weight-norm merge for Conv1d params: + - combines `parametrizations.weight.original0` (`g`) + `original1` (`v`) + - produces standard `.weight` + +Missing weight-norm pairs fail fast. + +## Quantization Guardrail + +`model_quant_predicate(...)` disables quantization for embedding modules to protect codebook fidelity. + +## Load/Save + +- `from_pretrained(path_or_repo, strict=True)` supports local path or HF repo. +- `save_config(path)` writes canonicalized config JSON. + +MOSS runtime post-load hooks look for embedded codec folders first (`audio_tokenizer`, `moss_audio_tokenizer`, `codec`) and then fall back to `OpenMOSS-Team/MOSS-Audio-Tokenizer`. + +## Validation Anchors + +Primary tests: + +- `mlx_audio/codec/tests/test_moss_audio_tokenizer.py` +- `mlx_audio/codec/tests/test_moss_audio_tokenizer_config_contracts.py` + +These cover layout normalization, prefix decode semantics, tie-resolution edge cases, sanitize behavior, and streaming decode constraints. diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/__init__.py b/mlx_audio/codec/models/moss_audio_tokenizer/__init__.py new file mode 100644 index 000000000..96310d464 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/__init__.py @@ -0,0 +1,23 @@ +from .config import ( + MossAudioTokenizerConfig, + MossAudioTokenizerModuleConfig, + MossAudioTokenizerQuantizerConfig, + load_moss_audio_tokenizer_config, +) +from .model import ( + MossAudioTokenizer, + MossAudioTokenizerDecoderOutput, + MossAudioTokenizerEncoderOutput, + MossAudioTokenizerOutput, +) + +__all__ = [ + "MossAudioTokenizer", + "MossAudioTokenizerConfig", + "MossAudioTokenizerDecoderOutput", + "MossAudioTokenizerEncoderOutput", + "MossAudioTokenizerModuleConfig", + "MossAudioTokenizerOutput", + "MossAudioTokenizerQuantizerConfig", + "load_moss_audio_tokenizer_config", +] diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/config.py b/mlx_audio/codec/models/moss_audio_tokenizer/config.py new file mode 100644 index 000000000..7fbdafe5c --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/config.py @@ -0,0 +1,206 @@ +"""Configuration helpers for the MOSS audio tokenizer MLX port.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict, List, Optional + +LEGACY_MODEL_TYPES = {"speech_tokenizer", "moss-audio-tokenizer"} +CANONICAL_MODEL_TYPE = "moss_audio_tokenizer" + + +@dataclass(frozen=True) +class MossAudioTokenizerModuleConfig: + module_type: str + patch_size: Optional[int] = None + input_dimension: Optional[int] = None + output_dimension: Optional[int] = None + d_model: Optional[int] = None + num_heads: Optional[int] = None + num_layers: Optional[int] = None + dim_feedforward: Optional[int] = None + causal: Optional[bool] = None + norm: Optional[str] = None + positional_embedding: Optional[str] = None + max_period: Optional[float] = None + gating: Optional[str] = None + layer_scale: Optional[float] = None + conv_layout: Optional[bool] = None + + @classmethod + def from_dict(cls, data: Dict[str, Any]) -> "MossAudioTokenizerModuleConfig": + return cls( + module_type=str(data["module_type"]), + patch_size=data.get("patch_size"), + input_dimension=data.get("input_dimension"), + output_dimension=data.get("output_dimension"), + d_model=data.get("d_model"), + num_heads=data.get("num_heads"), + num_layers=data.get("num_layers"), + dim_feedforward=data.get("dim_feedforward"), + causal=data.get("causal"), + norm=data.get("norm"), + positional_embedding=data.get("positional_embedding"), + max_period=data.get("max_period"), + gating=data.get("gating"), + layer_scale=data.get("layer_scale"), + conv_layout=data.get("conv_layout"), + ) + + def to_dict(self) -> Dict[str, Any]: + payload: Dict[str, Any] = {"module_type": self.module_type} + for field_name in [ + "patch_size", + "input_dimension", + "output_dimension", + "d_model", + "num_heads", + "num_layers", + "dim_feedforward", + "causal", + "norm", + "positional_embedding", + "max_period", + "gating", + "layer_scale", + "conv_layout", + ]: + value = getattr(self, field_name) + if value is not None: + payload[field_name] = value + return payload + + +@dataclass(frozen=True) +class MossAudioTokenizerQuantizerConfig: + input_dim: int + rvq_dim: int + output_dim: int + num_quantizers: int + codebook_size: int + codebook_dim: int + quantizer_type: str + + @classmethod + def from_dict(cls, data: Dict[str, Any]) -> "MossAudioTokenizerQuantizerConfig": + return cls( + input_dim=int(data["input_dim"]), + rvq_dim=int(data["rvq_dim"]), + output_dim=int(data["output_dim"]), + num_quantizers=int(data["num_quantizers"]), + codebook_size=int(data["codebook_size"]), + codebook_dim=int(data["codebook_dim"]), + quantizer_type=str(data["quantizer_type"]), + ) + + def to_dict(self) -> Dict[str, Any]: + return { + "input_dim": self.input_dim, + "rvq_dim": self.rvq_dim, + "output_dim": self.output_dim, + "num_quantizers": self.num_quantizers, + "codebook_size": self.codebook_size, + "codebook_dim": self.codebook_dim, + "quantizer_type": self.quantizer_type, + } + + +@dataclass(frozen=True) +class MossAudioTokenizerConfig: + model_type: str + sampling_rate: int + downsample_rate: int + causal_transformer_context_duration: float + encoder_modules: List[MossAudioTokenizerModuleConfig] + decoder_modules: List[MossAudioTokenizerModuleConfig] + quantizer: MossAudioTokenizerQuantizerConfig + + @property + def frame_rate(self) -> float: + return self.sampling_rate / self.downsample_rate + + @property + def encoder_patch_product(self) -> int: + return _patch_product(self.encoder_modules) + + @property + def decoder_patch_product(self) -> int: + return _patch_product(self.decoder_modules) + + def patch_alignment_is_valid(self) -> bool: + # Encoder and decoder patch products should both match configured downsample. + return ( + self.encoder_patch_product == self.downsample_rate + and self.decoder_patch_product == self.downsample_rate + ) + + @classmethod + def from_dict(cls, data: Dict[str, Any]) -> "MossAudioTokenizerConfig": + model_type = str(data.get("model_type", CANONICAL_MODEL_TYPE)) + if model_type in LEGACY_MODEL_TYPES: + model_type = CANONICAL_MODEL_TYPE + + encoder_payload = data.get("encoder_kwargs", data.get("encoder_modules", [])) + decoder_payload = data.get("decoder_kwargs", data.get("decoder_modules", [])) + encoder_modules = [ + MossAudioTokenizerModuleConfig.from_dict(item) for item in encoder_payload + ] + decoder_modules = [ + MossAudioTokenizerModuleConfig.from_dict(item) for item in decoder_payload + ] + quantizer_payload = data.get("quantizer_kwargs", data.get("quantizer")) + if quantizer_payload is None: + raise ValueError("Missing quantizer_kwargs/quantizer in config payload") + quantizer = MossAudioTokenizerQuantizerConfig.from_dict(quantizer_payload) + + return cls( + model_type=model_type, + sampling_rate=int(data["sampling_rate"]), + downsample_rate=int(data["downsample_rate"]), + causal_transformer_context_duration=float( + data["causal_transformer_context_duration"] + ), + encoder_modules=encoder_modules, + decoder_modules=decoder_modules, + quantizer=quantizer, + ) + + def to_dict(self) -> Dict[str, Any]: + return { + "model_type": self.model_type, + "sampling_rate": self.sampling_rate, + "downsample_rate": self.downsample_rate, + "causal_transformer_context_duration": self.causal_transformer_context_duration, + "encoder_kwargs": [module.to_dict() for module in self.encoder_modules], + "decoder_kwargs": [module.to_dict() for module in self.decoder_modules], + "quantizer_kwargs": self.quantizer.to_dict(), + "quantizer_type": self.quantizer.quantizer_type, + } + + +def _patch_product(modules: List[MossAudioTokenizerModuleConfig]) -> int: + product = 1 + for module in modules: + if module.module_type == "PatchedPretransform": + if module.patch_size is None: + raise ValueError("PatchedPretransform module is missing patch_size") + product *= int(module.patch_size) + return product + + +def load_moss_audio_tokenizer_config(path: str | Path) -> MossAudioTokenizerConfig: + config_path = Path(path) + data = json.loads(config_path.read_text(encoding="utf-8")) + return MossAudioTokenizerConfig.from_dict(data) + + +__all__ = [ + "CANONICAL_MODEL_TYPE", + "LEGACY_MODEL_TYPES", + "MossAudioTokenizerConfig", + "MossAudioTokenizerModuleConfig", + "MossAudioTokenizerQuantizerConfig", + "load_moss_audio_tokenizer_config", +] diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/decoder.py b/mlx_audio/codec/models/moss_audio_tokenizer/decoder.py new file mode 100644 index 000000000..66da51ee4 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/decoder.py @@ -0,0 +1,48 @@ +"""Decoder module construction for the MOSS audio tokenizer.""" + +from __future__ import annotations + +import mlx.nn as nn + +from .config import MossAudioTokenizerConfig +from .modules import ( + MossAudioTokenizerPatchedPretransform, + MossAudioTokenizerProjectedTransformer, + MossAudioTokenizerTransformerConfig, +) + + +def build_moss_audio_tokenizer_decoder_modules( + config: MossAudioTokenizerConfig, +) -> list[nn.Module]: + current_frame_rate = float(config.sampling_rate) / float(config.downsample_rate) + modules: list[nn.Module] = [] + + for module_config in config.decoder_modules: + if module_config.module_type == "PatchedPretransform": + if module_config.patch_size is None: + raise ValueError("PatchedPretransform in decoder is missing patch_size") + module = MossAudioTokenizerPatchedPretransform( + patch_size=int(module_config.patch_size), + is_downsample=False, + module_type=module_config.module_type, + ) + elif module_config.module_type == "Transformer": + transformer_config = MossAudioTokenizerTransformerConfig.from_module_config( + module_config, + context=int( + current_frame_rate * config.causal_transformer_context_duration + ), + ) + module = MossAudioTokenizerProjectedTransformer( + transformer_config, module_type=module_config.module_type + ) + else: + raise ValueError( + f"Unsupported decoder module_type: {module_config.module_type}" + ) + + modules.append(module) + current_frame_rate *= int(getattr(module, "downsample_ratio", 1)) + + return modules diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/encoder.py b/mlx_audio/codec/models/moss_audio_tokenizer/encoder.py new file mode 100644 index 000000000..06ac8fd67 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/encoder.py @@ -0,0 +1,48 @@ +"""Encoder module construction for the MOSS audio tokenizer.""" + +from __future__ import annotations + +import mlx.nn as nn + +from .config import MossAudioTokenizerConfig +from .modules import ( + MossAudioTokenizerPatchedPretransform, + MossAudioTokenizerProjectedTransformer, + MossAudioTokenizerTransformerConfig, +) + + +def build_moss_audio_tokenizer_encoder_modules( + config: MossAudioTokenizerConfig, +) -> list[nn.Module]: + current_frame_rate = float(config.sampling_rate) + modules: list[nn.Module] = [] + + for module_config in config.encoder_modules: + if module_config.module_type == "PatchedPretransform": + if module_config.patch_size is None: + raise ValueError("PatchedPretransform in encoder is missing patch_size") + module = MossAudioTokenizerPatchedPretransform( + patch_size=int(module_config.patch_size), + is_downsample=True, + module_type=module_config.module_type, + ) + elif module_config.module_type == "Transformer": + transformer_config = MossAudioTokenizerTransformerConfig.from_module_config( + module_config, + context=int( + current_frame_rate * config.causal_transformer_context_duration + ), + ) + module = MossAudioTokenizerProjectedTransformer( + transformer_config, module_type=module_config.module_type + ) + else: + raise ValueError( + f"Unsupported encoder module_type: {module_config.module_type}" + ) + + modules.append(module) + current_frame_rate /= int(getattr(module, "downsample_ratio", 1)) + + return modules diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/model.py b/mlx_audio/codec/models/moss_audio_tokenizer/model.py new file mode 100644 index 000000000..e85784951 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/model.py @@ -0,0 +1,861 @@ +"""MLX implementation of the MOSS audio tokenizer codec.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, Optional, Union + +import mlx.core as mx +import mlx.nn as nn +from huggingface_hub import snapshot_download +from mlx.utils import tree_flatten + +from .config import MossAudioTokenizerConfig, load_moss_audio_tokenizer_config +from .decoder import build_moss_audio_tokenizer_decoder_modules +from .encoder import build_moss_audio_tokenizer_encoder_modules +from .quantizer import ( + MossAudioTokenizerResidualLFQ, + MossAudioTokenizerResidualVQ, + build_moss_audio_tokenizer_quantizer, +) + + +@dataclass +class MossAudioTokenizerEncoderOutput: + audio_codes: Optional[mx.array] = None + audio_codes_lengths: Optional[mx.array] = None + encoder_hidden_states: Optional[mx.array] = None + + +@dataclass +class MossAudioTokenizerDecoderOutput: + audio: Optional[mx.array] = None + audio_lengths: Optional[mx.array] = None + + +@dataclass +class MossAudioTokenizerOutput: + audio: Optional[mx.array] = None + audio_lengths: Optional[mx.array] = None + audio_codes: Optional[mx.array] = None + audio_codes_lengths: Optional[mx.array] = None + + +class MossAudioTokenizer(nn.Module): + """Standalone codec used by the MOSS-TTS family.""" + + def __init__(self, config: MossAudioTokenizerConfig): + super().__init__() + self.config = config + self.sampling_rate = config.sampling_rate + self.downsample_rate = config.downsample_rate + self.causal_transformer_context_duration = ( + config.causal_transformer_context_duration + ) + + self.encoder = build_moss_audio_tokenizer_encoder_modules(config) + self.quantizer = build_moss_audio_tokenizer_quantizer(config.quantizer) + self.decoder = build_moss_audio_tokenizer_decoder_modules(config) + + @property + def frame_rate(self) -> float: + return self.sampling_rate / self.downsample_rate + + def _build_caches(self, modules: list[nn.Module]) -> list[Optional[list]]: + caches: list[Optional[list]] = [] + for module in modules: + make_cache = getattr(module, "make_cache", None) + caches.append(make_cache() if callable(make_cache) else None) + return caches + + def _ensure_audio_layout(self, input_values: mx.array) -> mx.array: + if input_values.ndim == 1: + input_values = input_values[None, None, :] + elif input_values.ndim == 2: + input_values = input_values[:, None, :] + elif input_values.ndim != 3: + raise ValueError( + f"Expected input_values with 1/2/3 dims, got shape={input_values.shape}" + ) + return input_values + + def _resolve_requested_quantizers(self, num_quantizers: Optional[int]) -> int: + configured_quantizers = self.config.quantizer.num_quantizers + if num_quantizers is None: + return configured_quantizers + if num_quantizers <= 0: + raise ValueError("num_quantizers must be > 0 when provided.") + if num_quantizers > configured_quantizers: + raise ValueError( + f"num_quantizers ({num_quantizers}) must be <= configured " + f"quantizer count ({configured_quantizers})." + ) + return int(num_quantizers) + + def _normalize_decode_audio_codes( + self, + audio_codes: mx.array, + *, + num_quantizers: Optional[int], + ) -> mx.array: + configured_quantizers = self.config.quantizer.num_quantizers + requested_quantizers = self._resolve_requested_quantizers(num_quantizers) + + if audio_codes.ndim not in {2, 3}: + raise ValueError( + f"Expected audio_codes with 2 or 3 dims, got shape={audio_codes.shape}" + ) + + candidates: list[tuple[str, int, mx.array]] = [] + + def register_candidate( + orientation: str, + quantizer_count: int, + normalized: mx.array, + ) -> None: + candidates.append((orientation, quantizer_count, normalized)) + + if audio_codes.ndim == 2: + if audio_codes.shape[0] == configured_quantizers: + register_candidate( + "NQ-first", + configured_quantizers, + audio_codes[:, None, :], + ) + if ( + requested_quantizers != configured_quantizers + and audio_codes.shape[0] == requested_quantizers + ): + register_candidate( + "NQ-first", + requested_quantizers, + audio_codes[:, None, :], + ) + if audio_codes.shape[1] == configured_quantizers: + register_candidate( + "NQ-last", + configured_quantizers, + audio_codes.transpose(1, 0)[:, None, :], + ) + if ( + requested_quantizers != configured_quantizers + and audio_codes.shape[1] == requested_quantizers + ): + register_candidate( + "NQ-last", + requested_quantizers, + audio_codes.transpose(1, 0)[:, None, :], + ) + else: + if audio_codes.shape[0] == configured_quantizers: + register_candidate("NQ-first", configured_quantizers, audio_codes) + if ( + requested_quantizers != configured_quantizers + and audio_codes.shape[0] == requested_quantizers + ): + register_candidate("NQ-first", requested_quantizers, audio_codes) + if audio_codes.shape[-1] == configured_quantizers: + register_candidate( + "NQ-last", + configured_quantizers, + audio_codes.transpose(2, 0, 1), + ) + if ( + requested_quantizers != configured_quantizers + and audio_codes.shape[-1] == requested_quantizers + ): + register_candidate( + "NQ-last", + requested_quantizers, + audio_codes.transpose(2, 0, 1), + ) + + if not candidates: + shape_contract = ( + "Expected (NQ, T)/(T, NQ) or (NQ, B, T)/(B, T, NQ) with " + f"NQ in {{{configured_quantizers}" + + ( + f", {requested_quantizers}" + if requested_quantizers != configured_quantizers + else "" + ) + + "}." + ) + raise ValueError( + f"Unrecognized audio_codes layout for shape={audio_codes.shape}. " + f"{shape_contract}" + ) + + if len(candidates) > 1: + if ( + num_quantizers is not None + and requested_quantizers != configured_quantizers + ): + # Only prefix requests should override canonical tie resolution. + # For no-op requests (requested == configured), preserve implicit + # canonical behavior so explicit num_quantizers does not transpose + # canonical decode inputs. + requested_candidates = [ + candidate + for candidate in candidates + if candidate[1] == requested_quantizers + ] + if len(requested_candidates) == 1: + candidates = requested_candidates + elif len(requested_candidates) > 1: + requested_nq_first = [ + candidate + for candidate in requested_candidates + if candidate[0] == "NQ-first" + ] + requested_nq_last = [ + candidate + for candidate in requested_candidates + if candidate[0] == "NQ-last" + ] + + # Explicit square ties are ambiguous by shape alone. Keep canonical + # NQ-first orientation for these cases so canonical decode inputs + # do not get transposed under explicit num_quantizers requests. + is_square_requested_tie = False + is_canonical_prefix_tie = False + if audio_codes.ndim == 2: + is_square_requested_tie = ( + int(audio_codes.shape[0]) == requested_quantizers + and int(audio_codes.shape[1]) == requested_quantizers + ) + else: + is_square_requested_tie = ( + int(audio_codes.shape[0]) == requested_quantizers + and int(audio_codes.shape[1]) == 1 + and int(audio_codes.shape[-1]) == requested_quantizers + ) or ( + int(audio_codes.shape[0]) == 1 + and int(audio_codes.shape[1]) == requested_quantizers + and int(audio_codes.shape[-1]) == requested_quantizers + ) + # True-prefix 3D ties can still be canonical `(NQ, B, T)` inputs + # when both first/last dimensions match the requested prefix + # count. Preserve NQ-first in this family so we do not swap + # batch/time semantics into `(NQ, T, B)`. + is_canonical_prefix_tie = ( + int(audio_codes.shape[0]) == requested_quantizers + and int(audio_codes.shape[-1]) == requested_quantizers + and int(audio_codes.shape[1]) > 1 + ) + + if (is_square_requested_tie or is_canonical_prefix_tie) and len( + requested_nq_first + ) == 1: + candidates = requested_nq_first + elif len(requested_nq_last) == 1: + candidates = requested_nq_last + else: + candidates = requested_candidates + if len(candidates) > 1: + # Fallback tie-break: prefer canonical layout so encode() -> decode() + # round-trips remain valid when T == NQ. + canonical_candidates = [ + candidate for candidate in candidates if candidate[0] == "NQ-first" + ] + if len(canonical_candidates) == 1: + candidates = canonical_candidates + else: + candidate_labels = [ + f"{orientation}[nq={nq}]" for orientation, nq, _ in candidates + ] + raise ValueError( + "Ambiguous audio_codes layout. Multiple interpretations are valid for " + f"shape={audio_codes.shape}: {candidate_labels}. " + "Use a canonical layout to disambiguate." + ) + + _, source_quantizers, normalized = candidates[0] + if requested_quantizers > source_quantizers: + raise ValueError( + f"num_quantizers ({requested_quantizers}) must be <= decoded source " + f"quantizer count ({source_quantizers})." + ) + if requested_quantizers < source_quantizers: + normalized = normalized[:requested_quantizers] + return normalized + + def _encode_frame( + self, + input_values: mx.array, + input_lengths: Optional[mx.array] = None, + n_quantizers: Optional[int] = None, + caches: Optional[list[Optional[list]]] = None, + ) -> MossAudioTokenizerEncoderOutput: + input_values = self._ensure_audio_layout(input_values) + batch_size, _, time_steps = input_values.shape + + if input_lengths is None: + input_lengths = mx.full((batch_size,), time_steps, dtype=mx.int32) + else: + input_lengths = input_lengths.astype(mx.int32) + + if time_steps % self.downsample_rate != 0: + pad_length = self.downsample_rate - (time_steps % self.downsample_rate) + input_values = mx.pad(input_values, [(0, 0), (0, 0), (0, pad_length)]) + + hidden = input_values + hidden_lengths = input_lengths + for idx, module in enumerate(self.encoder): + cache = caches[idx] if caches is not None else None + hidden, hidden_lengths = module(hidden, hidden_lengths, cache=cache) + + quantizer = self.quantizer + if not isinstance( + quantizer, (MossAudioTokenizerResidualLFQ, MossAudioTokenizerResidualVQ) + ): + raise TypeError(f"Unsupported quantizer type: {type(quantizer)}") + _, audio_codes, audio_codes_lengths = quantizer( + hidden, hidden_lengths, n_quantizers + ) + + return MossAudioTokenizerEncoderOutput( + audio_codes=audio_codes.astype(mx.int32), + audio_codes_lengths=audio_codes_lengths.astype(mx.int32), + encoder_hidden_states=hidden, + ) + + def _decode_frame( + self, + audio_codes: mx.array, + audio_codes_lengths: Optional[mx.array] = None, + caches: Optional[list[Optional[list]]] = None, + ) -> MossAudioTokenizerDecoderOutput: + if audio_codes.ndim != 3: + raise ValueError( + "Expected canonical audio_codes layout (NQ, B, T) in _decode_frame, " + f"got shape={audio_codes.shape}." + ) + audio_codes = audio_codes.astype(mx.int32) + _, batch_size, time_steps = audio_codes.shape + + if audio_codes_lengths is None: + audio_codes_lengths = mx.full((batch_size,), time_steps, dtype=mx.int32) + else: + audio_codes_lengths = audio_codes_lengths.astype(mx.int32) + + quantizer = self.quantizer + if not isinstance( + quantizer, (MossAudioTokenizerResidualLFQ, MossAudioTokenizerResidualVQ) + ): + raise TypeError(f"Unsupported quantizer type: {type(quantizer)}") + decoded = quantizer.decode_codes(audio_codes) + + audio = decoded + audio_lengths = audio_codes_lengths + for idx, module in enumerate(self.decoder): + cache = caches[idx] if caches is not None else None + audio, audio_lengths = module(audio, audio_lengths, cache=cache) + + return MossAudioTokenizerDecoderOutput( + audio=audio, + audio_lengths=audio_lengths.astype(mx.int32), + ) + + def encode( + self, + input_values: mx.array, + padding_mask: Optional[mx.array] = None, + num_quantizers: Optional[int] = None, + return_dict: bool = True, + chunk_duration: Optional[float] = None, + ) -> MossAudioTokenizerEncoderOutput | tuple[mx.array, mx.array]: + input_values = self._ensure_audio_layout(input_values) + batch_size, _, time_steps = input_values.shape + + if padding_mask is not None: + input_lengths = mx.sum(padding_mask.astype(mx.int32), axis=-1).astype( + mx.int32 + ) + else: + input_lengths = mx.full((batch_size,), time_steps, dtype=mx.int32) + + if chunk_duration is None: + output = self._encode_frame( + input_values, input_lengths, n_quantizers=num_quantizers + ) + else: + if chunk_duration <= 0: + raise ValueError("chunk_duration must be > 0 when provided.") + if chunk_duration > self.causal_transformer_context_duration: + raise ValueError( + "chunk_duration must be <= " + f"{self.causal_transformer_context_duration}, got {chunk_duration}." + ) + if batch_size != 1: + raise ValueError( + "Streaming encode currently only supports batch_size=1." + ) + + chunk_length = int(round(chunk_duration * self.sampling_rate)) + if chunk_length <= 0: + raise ValueError( + "chunk_duration is too small and results in chunk_length <= 0." + ) + if chunk_length % self.downsample_rate != 0: + raise ValueError( + "chunk_duration * sampling_rate must be divisible by downsample_rate. " + f"Got chunk_length={chunk_length}, downsample_rate={self.downsample_rate}." + ) + + input_length = int(input_lengths[0]) + if input_length <= chunk_length: + output = self._encode_frame( + input_values[..., :input_length], + input_lengths, + n_quantizers=num_quantizers, + ) + else: + caches = self._build_caches(self.encoder) + codes_chunks = [] + hidden_chunks = [] + for start_idx in range(0, input_length, chunk_length): + input_length_i = min(chunk_length, input_length - start_idx) + if input_length_i <= 0: + break + input_values_i = input_values[ + ..., start_idx : start_idx + input_length_i + ] + input_lengths_i = mx.array([input_length_i], dtype=mx.int32) + chunk_output = self._encode_frame( + input_values_i, + input_lengths_i, + n_quantizers=num_quantizers, + caches=caches, + ) + if ( + chunk_output.audio_codes is None + or chunk_output.audio_codes_lengths is None + ): + raise RuntimeError( + "Internal error: _encode_frame returned empty audio codes." + ) + if chunk_output.encoder_hidden_states is None: + raise RuntimeError( + "Internal error: _encode_frame returned empty hidden states." + ) + length_i = int(chunk_output.audio_codes_lengths[0]) + codes_chunks.append(chunk_output.audio_codes[:, :, :length_i]) + hidden_chunks.append( + chunk_output.encoder_hidden_states[:, :, :length_i] + ) + + audio_codes = mx.concatenate(codes_chunks, axis=-1) + hidden_states = mx.concatenate(hidden_chunks, axis=-1) + output = MossAudioTokenizerEncoderOutput( + audio_codes=audio_codes, + audio_codes_lengths=mx.array( + [audio_codes.shape[-1]], dtype=mx.int32 + ), + encoder_hidden_states=hidden_states, + ) + + if not return_dict: + if output.audio_codes is None or output.audio_codes_lengths is None: + raise RuntimeError("encode() produced empty outputs.") + return output.audio_codes, output.audio_codes_lengths + return output + + def decode( + self, + audio_codes: mx.array, + padding_mask: Optional[mx.array] = None, + return_dict: bool = True, + chunk_duration: Optional[float] = None, + num_quantizers: Optional[int] = None, + ) -> MossAudioTokenizerDecoderOutput | tuple[mx.array]: + audio_codes = self._normalize_decode_audio_codes( + audio_codes, + num_quantizers=num_quantizers, + ).astype(mx.int32) + + _, batch_size, time_steps = audio_codes.shape + if padding_mask is not None: + codes_lengths = mx.sum(padding_mask.astype(mx.int32), axis=-1).astype( + mx.int32 + ) + else: + codes_lengths = mx.full((batch_size,), time_steps, dtype=mx.int32) + + if chunk_duration is None: + output = self._decode_frame(audio_codes, codes_lengths) + else: + if chunk_duration <= 0: + raise ValueError("chunk_duration must be > 0 when provided.") + if chunk_duration > self.causal_transformer_context_duration: + raise ValueError( + "chunk_duration must be <= " + f"{self.causal_transformer_context_duration}, got {chunk_duration}." + ) + if batch_size != 1: + raise ValueError( + "Streaming decode currently only supports batch_size=1." + ) + + chunk_length = int(round(chunk_duration * self.sampling_rate)) + if chunk_length <= 0: + raise ValueError( + "chunk_duration is too small and results in chunk_length <= 0." + ) + if chunk_length % self.downsample_rate != 0: + raise ValueError( + "chunk_duration * sampling_rate must be divisible by downsample_rate. " + f"Got chunk_length={chunk_length}, downsample_rate={self.downsample_rate}." + ) + chunk_frame_length = chunk_length // self.downsample_rate + + codes_length = int(codes_lengths[0]) + if codes_length <= chunk_frame_length: + output = self._decode_frame( + audio_codes[..., :codes_length], codes_lengths + ) + else: + caches = self._build_caches(self.decoder) + wav_chunks = [] + for start_idx in range(0, codes_length, chunk_frame_length): + codes_length_i = min(chunk_frame_length, codes_length - start_idx) + if codes_length_i <= 0: + break + codes_i = audio_codes[:, :, start_idx : start_idx + codes_length_i] + codes_lengths_i = mx.array([codes_length_i], dtype=mx.int32) + chunk_output = self._decode_frame( + codes_i, codes_lengths_i, caches=caches + ) + if chunk_output.audio is None or chunk_output.audio_lengths is None: + raise RuntimeError( + "Internal error: _decode_frame returned empty audio." + ) + wav_chunk = chunk_output.audio[ + :, :, : int(chunk_output.audio_lengths[0]) + ] + mx.eval(wav_chunk) + wav_chunks.append(wav_chunk) + mx.clear_cache() + wav = mx.concatenate(wav_chunks, axis=-1) + output = MossAudioTokenizerDecoderOutput( + audio=wav, + audio_lengths=mx.array([wav.shape[-1]], dtype=mx.int32), + ) + + if not return_dict: + if output.audio is None: + raise RuntimeError("decode() produced empty audio.") + return (output.audio,) + return output + + def batch_encode( + self, + wav_list: list[mx.array], + num_quantizers: Optional[int] = None, + ) -> MossAudioTokenizerEncoderOutput: + if not wav_list: + raise ValueError("wav_list must contain at least one waveform.") + + normalized = [] + for wav in wav_list: + if wav.ndim == 2: + if wav.shape[0] != 1: + raise ValueError( + "Expected 2D waveform shape (1, T), got shape=" f"{wav.shape}." + ) + wav = wav.squeeze(0) + if wav.ndim != 1: + raise ValueError( + f"Each waveform in wav_list must be 1D or (1, T), got {wav.shape}." + ) + normalized.append(wav) + + batch_size = len(normalized) + max_length = max(int(w.shape[-1]) for w in normalized) + input_values = mx.zeros((batch_size, 1, max_length), dtype=mx.float32) + input_lengths = mx.zeros((batch_size,), dtype=mx.int32) + + for idx, wav in enumerate(normalized): + length_i = int(wav.shape[-1]) + input_values[idx, 0, :length_i] = wav.astype(mx.float32) + input_lengths[idx] = length_i + + return self._encode_frame( + input_values, + input_lengths, + n_quantizers=num_quantizers, + ) + + def batch_decode( + self, + codes_list: list[mx.array], + num_quantizers: Optional[int] = None, + ) -> MossAudioTokenizerDecoderOutput: + if not codes_list: + raise ValueError("codes_list must contain at least one code tensor.") + + normalized = [] + for codes in codes_list: + normalized_codes = self._normalize_decode_audio_codes( + codes, + num_quantizers=num_quantizers, + ) + if int(normalized_codes.shape[1]) != 1: + raise ValueError( + "batch_decode() expects each codes_list entry to resolve to " + f"batch_size=1, got batch_size={int(normalized_codes.shape[1])} " + f"for shape={codes.shape}. Use decode() for batched code tensors." + ) + normalized.append(normalized_codes.squeeze(1)) + target_quantizers = int(normalized[0].shape[0]) + if any(int(c.shape[0]) != target_quantizers for c in normalized): + raise ValueError( + "All elements in codes_list must resolve to the same quantizer count." + ) + + max_length = max(int(c.shape[-1]) for c in normalized) + batch_size = len(normalized) + audio_codes = mx.zeros( + (target_quantizers, batch_size, max_length), + dtype=mx.int32, + ) + audio_codes_lengths = mx.zeros((batch_size,), dtype=mx.int32) + + for idx, codes in enumerate(normalized): + time_steps = int(codes.shape[-1]) + audio_codes[:, idx, :time_steps] = codes + audio_codes_lengths[idx] = time_steps + + return self._decode_frame(audio_codes, audio_codes_lengths) + + def streaming_decode( + self, + audio_codes: mx.array, + *, + chunk_tokens: int = 100, + num_quantizers: Optional[int] = None, + ): + if chunk_tokens <= 0: + raise ValueError("chunk_tokens must be > 0.") + + audio_codes = self._normalize_decode_audio_codes( + audio_codes, + num_quantizers=num_quantizers, + ).astype(mx.int32) + + _, batch_size, total_tokens = audio_codes.shape + if batch_size != 1: + raise ValueError("streaming_decode currently only supports batch_size=1.") + + caches = self._build_caches(self.decoder) + for start_idx in range(0, total_tokens, chunk_tokens): + end_idx = min(start_idx + chunk_tokens, total_tokens) + codes_chunk = audio_codes[:, :, start_idx:end_idx] + chunk_len = int(codes_chunk.shape[-1]) + output = self._decode_frame( + codes_chunk, + mx.array([chunk_len], dtype=mx.int32), + caches=caches, + ) + if output.audio is None or output.audio_lengths is None: + raise RuntimeError( + "Internal error: _decode_frame returned empty audio." + ) + wav_chunk = output.audio[:, :, : int(output.audio_lengths[0])] + mx.eval(wav_chunk) + yield wav_chunk + mx.clear_cache() + + def __call__( + self, + input_values: Optional[mx.array] = None, + padding_mask: Optional[mx.array] = None, + audio_codes: Optional[mx.array] = None, + num_quantizers: Optional[int] = None, + return_dict: bool = True, + ) -> ( + MossAudioTokenizerOutput + | tuple[Optional[mx.array], Optional[mx.array], Optional[mx.array]] + ): + output_audio_codes = None + output_audio_codes_lengths = None + output_audio = None + output_audio_lengths = None + decoded_from_encoded_codes = False + + if input_values is not None: + encode_output = self.encode( + input_values, + padding_mask=padding_mask, + num_quantizers=num_quantizers, + return_dict=True, + ) + if not isinstance(encode_output, MossAudioTokenizerEncoderOutput): + raise RuntimeError( + "Internal error: encode() returned unexpected output type." + ) + output_audio_codes = encode_output.audio_codes + output_audio_codes_lengths = encode_output.audio_codes_lengths + + if audio_codes is None: + audio_codes = output_audio_codes + decoded_from_encoded_codes = True + + if audio_codes is not None: + if decoded_from_encoded_codes and output_audio_codes_lengths is not None: + decode_output = self._decode_frame( + audio_codes, output_audio_codes_lengths + ) + else: + decode_output = self.decode( + audio_codes, + padding_mask=padding_mask, + return_dict=True, + num_quantizers=num_quantizers, + ) + if not isinstance(decode_output, MossAudioTokenizerDecoderOutput): + raise RuntimeError( + "Internal error: decode() returned unexpected output type." + ) + output_audio = decode_output.audio + output_audio_lengths = decode_output.audio_lengths + + if not return_dict: + return output_audio_codes, output_audio, output_audio_lengths + + return MossAudioTokenizerOutput( + audio=output_audio, + audio_lengths=output_audio_lengths, + audio_codes=output_audio_codes, + audio_codes_lengths=output_audio_codes_lengths, + ) + + def model_quant_predicate(self, path: str, module) -> bool: + # Protect codebook embeddings from quantization. + if isinstance(module, nn.Embedding): + return False + return True + + def _sanitize_chunk( + self, + weights: Dict[str, mx.array], + pending_weight_norm: Optional[Dict[str, Dict[str, mx.array]]] = None, + ) -> tuple[Dict[str, mx.array], Dict[str, Dict[str, mx.array]]]: + pending = {} if pending_weight_norm is None else dict(pending_weight_norm) + sanitized: Dict[str, mx.array] = {} + current_shapes = { + name: tuple(value.shape) for name, value in tree_flatten(self.parameters()) + } + + for key, value in weights.items(): + if key.endswith(".parametrizations.weight.original0"): + base = key[: -len(".parametrizations.weight.original0")] + entry = pending.setdefault(base, {}) + entry["g"] = value + continue + if key.endswith(".parametrizations.weight.original1"): + base = key[: -len(".parametrizations.weight.original1")] + entry = pending.setdefault(base, {}) + entry["v"] = value + continue + + new_key = key + if value.ndim == 3 and new_key in current_shapes: + if tuple(value.shape) != current_shapes[new_key]: + value = value.swapaxes(1, 2) + sanitized[new_key] = value + + resolved = [] + for base, parts in pending.items(): + if "g" not in parts or "v" not in parts: + continue + g = parts["g"].astype(mx.float32) + v = parts["v"].astype(mx.float32) + reduce_axes = tuple(range(1, v.ndim)) + norm = mx.sqrt(mx.sum(v**2, axis=reduce_axes, keepdims=True) + 1e-12) + merged = (g * v / norm).astype(parts["v"].dtype) + target_key = f"{base}.weight" + if merged.ndim == 3 and target_key in current_shapes: + if tuple(merged.shape) != current_shapes[target_key]: + merged = merged.swapaxes(1, 2) + sanitized[target_key] = merged + resolved.append(base) + + for base in resolved: + pending.pop(base, None) + + return sanitized, pending + + def sanitize(self, weights: Dict[str, mx.array]) -> Dict[str, mx.array]: + sanitized, pending = self._sanitize_chunk(weights, pending_weight_norm=None) + if pending: + missing = sorted(pending.keys()) + raise ValueError( + "Missing weight-norm pairs while sanitizing weights for keys: " + f"{missing}" + ) + return sanitized + + @classmethod + def from_pretrained( + cls, + path_or_repo: Union[str, Path], + *, + strict: bool = True, + ) -> "MossAudioTokenizer": + model_path = Path(path_or_repo) + if not model_path.exists(): + model_path = Path( + snapshot_download( + str(path_or_repo), + allow_patterns=["*.json", "*.safetensors"], + ) + ) + + config_path = model_path / "config.json" + if not config_path.exists(): + raise FileNotFoundError(f"Missing config.json at {model_path}") + config = load_moss_audio_tokenizer_config(config_path) + model = cls(config) + + weight_files = sorted(model_path.glob("*.safetensors")) + if not weight_files: + raise FileNotFoundError(f"No *.safetensors files found at {model_path}") + + loaded_keys = set() + pending: Dict[str, Dict[str, mx.array]] = {} + for weight_file in weight_files: + shard_weights = mx.load(weight_file.as_posix(), format="safetensors") + sanitized, pending = model._sanitize_chunk( + shard_weights, pending_weight_norm=pending + ) + if sanitized: + model.load_weights(list(sanitized.items()), strict=False) + loaded_keys.update(sanitized.keys()) + + if pending: + unresolved = sorted(pending.keys()) + raise ValueError( + "Unresolved weight-norm parameter groups across shards: " + f"{unresolved}" + ) + + if strict: + expected_keys = {name for name, _ in tree_flatten(model.parameters())} + missing = sorted(expected_keys - loaded_keys) + if missing: + raise ValueError( + "Strict load failed: missing parameter weights for keys: " + f"{missing[:20]}" + ("..." if len(missing) > 20 else "") + ) + + mx.eval(model.parameters()) + return model + + def save_config(self, path: Union[str, Path]) -> None: + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as handle: + json.dump(self.config.to_dict(), handle, indent=2, sort_keys=True) diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/modules.py b/mlx_audio/codec/models/moss_audio_tokenizer/modules.py new file mode 100644 index 000000000..22005b4d2 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/modules.py @@ -0,0 +1,424 @@ +"""Core modules for the MOSS audio tokenizer codec. + +This file intentionally mirrors upstream module naming where practical so +checkpoint weights can load with minimal remapping. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional, Sequence + +import mlx.core as mx +import mlx.nn as nn +from mlx_lm.models.cache import KVCache, RotatingKVCache + +from .config import MossAudioTokenizerModuleConfig + + +@dataclass(frozen=True) +class MossAudioTokenizerTransformerConfig: + input_dimension: int + output_dimension: int + d_model: int + num_heads: int + num_layers: int + dim_feedforward: int + causal: bool = True + norm: str = "layer_norm" + positional_embedding: str = "rope" + max_period: float = 10000.0 + gating: str = "none" + layer_scale: Optional[float] = None + conv_layout: bool = True + context: Optional[int] = None + + @classmethod + def from_module_config( + cls, + module_config: MossAudioTokenizerModuleConfig, + *, + context: Optional[int], + ) -> "MossAudioTokenizerTransformerConfig": + if module_config.module_type != "Transformer": + raise ValueError( + f"Expected Transformer module config, got {module_config.module_type}" + ) + required_int_fields = { + "input_dimension": module_config.input_dimension, + "output_dimension": module_config.output_dimension, + "d_model": module_config.d_model, + "num_heads": module_config.num_heads, + "num_layers": module_config.num_layers, + } + for field_name, value in required_int_fields.items(): + if value is None: + raise ValueError( + f"Transformer module config is missing {field_name}: {module_config}" + ) + + dim_feedforward = module_config.dim_feedforward + if dim_feedforward is None: + raise ValueError( + "Transformer module config is missing dim_feedforward: " + f"{module_config}" + ) + + return cls( + input_dimension=int(module_config.input_dimension), + output_dimension=int(module_config.output_dimension), + d_model=int(module_config.d_model), + num_heads=int(module_config.num_heads), + num_layers=int(module_config.num_layers), + dim_feedforward=int(dim_feedforward), + causal=bool( + module_config.causal if module_config.causal is not None else True + ), + norm=str( + module_config.norm if module_config.norm is not None else "layer_norm" + ), + positional_embedding=str( + module_config.positional_embedding + if module_config.positional_embedding is not None + else "rope" + ), + max_period=float( + module_config.max_period + if module_config.max_period is not None + else 10000 + ), + gating=str( + module_config.gating if module_config.gating is not None else "none" + ), + layer_scale=( + float(module_config.layer_scale) + if module_config.layer_scale is not None + else None + ), + conv_layout=bool( + module_config.conv_layout + if module_config.conv_layout is not None + else True + ), + context=context, + ) + + @property + def head_dim(self) -> int: + if self.d_model % self.num_heads != 0: + raise ValueError( + f"d_model ({self.d_model}) must be divisible by num_heads ({self.num_heads})" + ) + return self.d_model // self.num_heads + + +class MossAudioTokenizerLayerScale(nn.Module): + """Per-channel learned residual scaling.""" + + def __init__(self, channels: int, init: float): + super().__init__() + self.scale = mx.ones((channels,)) * init + + def __call__(self, x: mx.array) -> mx.array: + return x * self.scale + + +def _create_norm(norm_type: str, dim: int) -> nn.Module: + if norm_type == "layer_norm": + return nn.LayerNorm(dim, eps=1e-5) + if norm_type == "rms_norm": + return nn.RMSNorm(dim, eps=1e-8) + raise ValueError(f"Unsupported norm type: {norm_type}") + + +def _apply_weights_per_step( + modules: Sequence[nn.Module], + schedule: Optional[list[int]], + x: mx.array, + offset: int, +) -> mx.array: + if len(modules) == 1: + return modules[0](x) + + outputs = [] + _, time_steps, _ = x.shape + for step in range(time_steps): + module_index = step + offset + if schedule is not None: + if module_index < 0 or module_index >= len(schedule): + raise ValueError( + "weights_per_step_schedule is too short for " + f"module_index={module_index}." + ) + module_index = schedule[module_index] + if module_index < 0 or module_index >= len(modules): + raise ValueError( + f"module_index={module_index} is out of range for {len(modules)} modules." + ) + outputs.append(modules[module_index](x[:, step : step + 1])) + return mx.concatenate(outputs, axis=1) + + +class MossAudioTokenizerMultiheadAttention(nn.Module): + """Causal MHA with optional RoPE and bounded KV cache context.""" + + def __init__(self, config: MossAudioTokenizerTransformerConfig): + super().__init__() + self.embed_dim = config.d_model + self.num_heads = config.num_heads + self.head_dim = config.head_dim + self.scale = self.head_dim**-0.5 + self.causal = config.causal + self.context = config.context + self.weights_per_step_schedule: Optional[list[int]] = None + + # Kept as lists to match upstream checkpoint key names: + # self_attn.in_projs.0.weight and self_attn.out_projs.0.weight + self.in_projs = [nn.Linear(self.embed_dim, 3 * self.embed_dim, bias=False)] + self.out_projs = [nn.Linear(self.embed_dim, self.embed_dim, bias=False)] + + self.rope = None + if config.positional_embedding in {"rope", "sin_rope"}: + self.rope = nn.RoPE(self.head_dim, traditional=True, base=config.max_period) + + def __call__( + self, + x: mx.array, + cache: Optional[KVCache | RotatingKVCache] = None, + mask: Optional[mx.array] = None, + ) -> mx.array: + batch_size, time_steps, hidden_dim = x.shape + if hidden_dim != self.embed_dim: + raise ValueError(f"Expected hidden dim {self.embed_dim}, got {hidden_dim}") + + offset = 0 if cache is None else int(cache.offset) + projected = _apply_weights_per_step( + self.in_projs, + self.weights_per_step_schedule, + x, + offset, + ) + projected = projected.reshape( + batch_size, time_steps, 3, self.num_heads, self.head_dim + ).transpose(2, 0, 3, 1, 4) + + q = projected[0] + k = projected[1] + v = projected[2] + + if self.rope is not None: + q = self.rope(q, offset=offset) + k = self.rope(k, offset=offset) + + if cache is not None: + k, v = cache.update_and_fetch(k, v) + + if mask is None and self.causal: + key_steps = k.shape[2] + key_offset = 0 if cache is None else int(cache.offset) - key_steps + key_positions = mx.arange(key_steps, dtype=mx.int32) + key_offset + query_positions = mx.arange(time_steps, dtype=mx.int32) + offset + delta = query_positions[:, None] - key_positions[None, :] + allowed = (key_positions[None, :] >= 0) & (delta >= 0) + if self.context is not None: + allowed = allowed & (delta < self.context) + mask = mx.where(allowed, 0.0, -1e9).astype(x.dtype) + mask = mask[None, None, :, :] + + attended = mx.fast.scaled_dot_product_attention( + q, + k, + v, + scale=self.scale, + mask=mask, + ) + attended = attended.transpose(0, 2, 1, 3).reshape(batch_size, time_steps, -1) + return _apply_weights_per_step( + self.out_projs, + self.weights_per_step_schedule, + attended, + offset, + ) + + +class MossAudioTokenizerTransformerLayer(nn.Module): + """Transformer block aligned to the upstream module contract.""" + + def __init__(self, config: MossAudioTokenizerTransformerConfig): + super().__init__() + if config.gating != "none": + raise ValueError( + f"Unsupported gating mode for MOSS audio tokenizer: {config.gating}" + ) + + self.self_attn = MossAudioTokenizerMultiheadAttention(config) + self.norm1 = _create_norm(config.norm, config.d_model) + self.norm2 = _create_norm(config.norm, config.d_model) + self.linear1 = nn.Linear(config.d_model, config.dim_feedforward, bias=False) + self.linear2 = nn.Linear(config.dim_feedforward, config.d_model, bias=False) + + if config.layer_scale is None: + self.layer_scale_1 = nn.Identity() + self.layer_scale_2 = nn.Identity() + else: + self.layer_scale_1 = MossAudioTokenizerLayerScale( + config.d_model, config.layer_scale + ) + self.layer_scale_2 = MossAudioTokenizerLayerScale( + config.d_model, config.layer_scale + ) + + def __call__( + self, + x: mx.array, + cache: Optional[KVCache | RotatingKVCache] = None, + mask: Optional[mx.array] = None, + ) -> mx.array: + attn_update = self.self_attn(self.norm1(x), cache=cache, mask=mask) + x = x + self.layer_scale_1(attn_update) + mlp_update = self.linear2(nn.gelu_approx(self.linear1(self.norm2(x)))) + return x + self.layer_scale_2(mlp_update) + + +class MossAudioTokenizerTransformer(nn.Module): + """Stacked transformer with per-layer KV caches.""" + + def __init__(self, config: MossAudioTokenizerTransformerConfig): + super().__init__() + self.config = config + self.layers = [ + MossAudioTokenizerTransformerLayer(config) for _ in range(config.num_layers) + ] + + def make_cache(self) -> list[KVCache | RotatingKVCache]: + if self.config.context is None: + return [KVCache() for _ in self.layers] + return [ + RotatingKVCache(max_size=self.config.context, keep=0) for _ in self.layers + ] + + def __call__( + self, + x: mx.array, + cache: Optional[list[KVCache | RotatingKVCache]] = None, + mask: Optional[mx.array] = None, + ) -> mx.array: + if cache is None: + per_layer_cache: list[Optional[KVCache | RotatingKVCache]] = [ + None for _ in self.layers + ] + else: + if len(cache) != len(self.layers): + raise ValueError( + "Cache depth mismatch: expected " + f"{len(self.layers)} layers, got {len(cache)}." + ) + per_layer_cache = list(cache) + + for layer, layer_cache in zip(self.layers, per_layer_cache): + x = layer(x, cache=layer_cache, mask=mask) + return x + + +class MossAudioTokenizerProjectedTransformer(nn.Module): + """Transformer block with optional input/output projections.""" + + def __init__( + self, + config: MossAudioTokenizerTransformerConfig, + *, + module_type: str = "Transformer", + ): + super().__init__() + self.module_type = module_type + self.downsample_ratio = 1 + self.conv_layout = config.conv_layout + self.input_dimension = config.input_dimension + self.output_dimension = config.output_dimension + + if config.input_dimension == config.d_model: + self.input_proj = None + else: + self.input_proj = nn.Linear( + config.input_dimension, config.d_model, bias=False + ) + self.transformer = MossAudioTokenizerTransformer(config) + if config.output_dimension == config.d_model: + self.output_proj = None + else: + self.output_proj = nn.Linear( + config.d_model, config.output_dimension, bias=False + ) + + def make_cache(self) -> list[KVCache | RotatingKVCache]: + return self.transformer.make_cache() + + def __call__( + self, + x: mx.array, + input_lengths: mx.array, + cache: Optional[list[KVCache | RotatingKVCache]] = None, + mask: Optional[mx.array] = None, + ) -> tuple[mx.array, mx.array]: + if self.conv_layout: + x = x.swapaxes(1, 2) + if self.input_proj is not None: + x = self.input_proj(x) + x = self.transformer(x, cache=cache, mask=mask) + if self.output_proj is not None: + x = self.output_proj(x) + if self.conv_layout: + x = x.swapaxes(1, 2) + return x, input_lengths + + +class MossAudioTokenizerPatchedPretransform(nn.Module): + """Patch/unpatch module used for deterministic stride changes.""" + + def __init__( + self, + patch_size: int, + is_downsample: bool, + *, + module_type: str = "PatchedPretransform", + ): + super().__init__() + self.patch_size = int(patch_size) + self.downsample_ratio = self.patch_size + self.is_downsample = is_downsample + self.module_type = module_type + + def encode(self, x: mx.array, input_lengths: mx.array) -> tuple[mx.array, mx.array]: + batch_size, channels, _ = x.shape + patch = self.patch_size + x = ( + x.reshape(batch_size, channels, -1, patch) + .transpose(0, 1, 3, 2) + .reshape(batch_size, channels * patch, -1) + ) + output_lengths = input_lengths // patch + return x, output_lengths + + def decode(self, x: mx.array, input_lengths: mx.array) -> tuple[mx.array, mx.array]: + batch_size, channels_times_patch, length = x.shape + patch = self.patch_size + channels = channels_times_patch // patch + x = ( + x.reshape(batch_size, channels, patch, length) + .transpose(0, 1, 3, 2) + .reshape(batch_size, channels, length * patch) + ) + output_lengths = input_lengths * patch + return x, output_lengths + + def __call__( + self, + x: mx.array, + input_lengths: mx.array, + cache=None, + mask=None, + ) -> tuple[mx.array, mx.array]: + del cache, mask + if self.is_downsample: + return self.encode(x, input_lengths) + return self.decode(x, input_lengths) diff --git a/mlx_audio/codec/models/moss_audio_tokenizer/quantizer.py b/mlx_audio/codec/models/moss_audio_tokenizer/quantizer.py new file mode 100644 index 000000000..619a48dc3 --- /dev/null +++ b/mlx_audio/codec/models/moss_audio_tokenizer/quantizer.py @@ -0,0 +1,310 @@ +"""Quantizer modules for the MOSS audio tokenizer.""" + +from __future__ import annotations + +from typing import Optional + +import mlx.core as mx +import mlx.nn as nn + +from mlx_audio.codec.models.mimi.modules.conv import Conv1d + +from .config import MossAudioTokenizerQuantizerConfig + + +def _l2_normalize(x: mx.array, axis: int) -> mx.array: + norm = mx.sqrt(mx.sum(x.astype(mx.float32) ** 2, axis=axis, keepdims=True) + 1e-12) + return x / norm + + +def _mask_from_lengths(lengths: mx.array, max_time: int, dtype) -> mx.array: + mask = mx.arange(max_time, dtype=lengths.dtype)[None, :] < lengths[:, None] + return mask.astype(dtype)[:, None, :] + + +class MossAudioTokenizerVectorQuantize(nn.Module): + """Single RVQ codebook (inference-only).""" + + def __init__( + self, + input_dim: int, + codebook_size: int, + codebook_dim: int, + ): + super().__init__() + self.input_dim = input_dim + self.codebook_size = codebook_size + self.codebook_dim = codebook_dim + + if input_dim != codebook_dim: + self.in_proj = Conv1d(input_dim, codebook_dim, 1, bias=True) + self.out_proj = Conv1d(codebook_dim, input_dim, 1, bias=True) + else: + self.in_proj = None + self.out_proj = None + + self.codebook = nn.Embedding(codebook_size, codebook_dim) + + def decode_code(self, indices: mx.array) -> mx.array: + # (B, T) -> (B, D, T) + quantized = self.codebook(indices).swapaxes(1, 2).astype(mx.float32) + if self.out_proj is not None: + quantized = self.out_proj(quantized) + return quantized.astype(mx.float32) + + def __call__(self, z: mx.array) -> tuple[mx.array, mx.array, mx.array]: + z = z.astype(mx.float32) + z_e = self.in_proj(z) if self.in_proj is not None else z + encodings = z_e.swapaxes(1, 2).reshape(-1, z_e.shape[1]).astype(mx.float32) + codebook = self.codebook.weight.astype(mx.float32) + + distances = ( + mx.sum(encodings**2, axis=1, keepdims=True) + - 2.0 * (encodings @ codebook.transpose()) + + mx.sum(codebook**2, axis=1)[None, :] + ) + indices = mx.argmax(-distances, axis=1).reshape(z.shape[0], -1).astype(mx.int32) + z_q = self.decode_code(indices) + return z_q, indices, z_e.astype(mx.float32) + + +class MossAudioTokenizerLFQ(nn.Module): + """LFQ codebook used by ResidualLFQ.""" + + def __init__( + self, + input_dim: int, + codebook_size: int, + codebook_dim: int, + ): + super().__init__() + self.input_dim = input_dim + self.codebook_size = codebook_size + self.codebook_dim = codebook_dim + + if input_dim != codebook_dim: + self.in_proj = Conv1d(input_dim, codebook_dim, 1, bias=True) + self.out_proj = Conv1d(codebook_dim, input_dim, 1, bias=True) + else: + self.in_proj = None + self.out_proj = None + + self.codebook = nn.Embedding(codebook_size, codebook_dim) + + def decode_code_wo_out_proj(self, indices: mx.array) -> mx.array: + return self.codebook(indices).swapaxes(1, 2).astype(mx.float32) + + def decode_code(self, indices: mx.array) -> mx.array: + z_q = self.decode_code_wo_out_proj(indices).astype(mx.float32) + if self.out_proj is not None: + z_q = self.out_proj(z_q) + return z_q.astype(mx.float32) + + def _decode_latents(self, latents: mx.array) -> tuple[mx.array, mx.array]: + encodings = ( + latents.swapaxes(1, 2).reshape(-1, latents.shape[1]).astype(mx.float32) + ) + codebook = self.codebook.weight.astype(mx.float32) + + encodings = _l2_normalize(encodings, axis=-1) + codebook = _l2_normalize(codebook, axis=-1) + + distances = ( + mx.sum(encodings**2, axis=1, keepdims=True) + - 2.0 * (encodings @ codebook.transpose()) + + mx.sum(codebook**2, axis=1)[None, :] + ) + indices = ( + mx.argmax(-distances, axis=1).reshape(latents.shape[0], -1).astype(mx.int32) + ) + z_q = self.decode_code_wo_out_proj(indices) + return z_q.astype(mx.float32), indices + + def __call__(self, z: mx.array) -> tuple[mx.array, mx.array, mx.array]: + z = z.astype(mx.float32) + z_e = self.in_proj(z) if self.in_proj is not None else z + z_q, indices = self._decode_latents(z_e.astype(mx.float32)) + if self.out_proj is not None: + z_q = self.out_proj(z_q) + return z_q.astype(mx.float32), indices, z_e.astype(mx.float32) + + +class MossAudioTokenizerResidualVQ(nn.Module): + """Residual VQ stack.""" + + def __init__( + self, + input_dim: int, + rvq_dim: int, + output_dim: int, + num_quantizers: int, + codebook_size: int, + codebook_dim: int, + ): + super().__init__() + self.input_dim = input_dim + self.rvq_dim = rvq_dim + self.output_dim = output_dim + self.num_quantizers = num_quantizers + + self.input_proj = ( + Conv1d(input_dim, rvq_dim, 1, bias=True) if input_dim != rvq_dim else None + ) + self.output_proj = ( + Conv1d(rvq_dim, output_dim, 1, bias=True) if rvq_dim != output_dim else None + ) + self.quantizers = [ + MossAudioTokenizerVectorQuantize( + input_dim=rvq_dim, + codebook_size=codebook_size, + codebook_dim=codebook_dim, + ) + for _ in range(num_quantizers) + ] + + def __call__( + self, + z: mx.array, + input_length: mx.array, + n_quantizers: Optional[int] = None, + ) -> tuple[mx.array, mx.array, mx.array]: + z = self.input_proj(z) if self.input_proj is not None else z + z = z.astype(mx.float32) + batch_size, _, max_time = z.shape + mask = _mask_from_lengths(input_length, max_time, z.dtype) + + quantized_out = mx.zeros_like(z).astype(mx.float32) + residual = z.astype(mx.float32) + all_indices = [] + n_quantizers = int(n_quantizers or self.num_quantizers) + + for quantizer in self.quantizers[:n_quantizers]: + z_q_i, indices_i, _ = quantizer(residual * mask) + quantized_out = quantized_out + z_q_i * mask + residual = residual - z_q_i * mask + all_indices.append(indices_i.astype(mx.int32)) + + if all_indices: + audio_codes = mx.stack(all_indices, axis=0) + else: + audio_codes = mx.zeros((0, batch_size, max_time), dtype=mx.int32) + + if self.output_proj is not None: + quantized_out = self.output_proj(quantized_out) + return ( + quantized_out.astype(mx.float32), + audio_codes, + input_length.astype(mx.int32), + ) + + def decode_codes(self, codes: mx.array) -> mx.array: + nq, batch_size, time_steps = codes.shape + embeddings = mx.zeros((batch_size, self.rvq_dim, time_steps), dtype=mx.float32) + for i, quantizer in enumerate(self.quantizers[:nq]): + embeddings = embeddings + quantizer.decode_code(codes[i]).astype(mx.float32) + if self.output_proj is not None: + embeddings = self.output_proj(embeddings) + return embeddings.astype(mx.float32) + + +class MossAudioTokenizerResidualLFQ(nn.Module): + """Residual LFQ stack.""" + + def __init__( + self, + input_dim: int, + rvq_dim: int, + output_dim: int, + num_quantizers: int, + codebook_size: int, + codebook_dim: int, + ): + super().__init__() + self.input_dim = input_dim + self.rvq_dim = rvq_dim + self.output_dim = output_dim + self.num_quantizers = num_quantizers + + self.input_proj = ( + Conv1d(input_dim, rvq_dim, 1, bias=True) if input_dim != rvq_dim else None + ) + self.output_proj = ( + Conv1d(rvq_dim, output_dim, 1, bias=True) if rvq_dim != output_dim else None + ) + self.quantizers = [ + MossAudioTokenizerLFQ( + input_dim=rvq_dim, + codebook_size=codebook_size, + codebook_dim=codebook_dim, + ) + for _ in range(num_quantizers) + ] + + def __call__( + self, + z: mx.array, + input_length: mx.array, + n_quantizers: Optional[int] = None, + ) -> tuple[mx.array, mx.array, mx.array]: + z = self.input_proj(z) if self.input_proj is not None else z + z = z.astype(mx.float32) + batch_size, _, max_time = z.shape + mask = _mask_from_lengths(input_length, max_time, z.dtype) + + quantized_out = mx.zeros_like(z).astype(mx.float32) + residual = z.astype(mx.float32) + all_indices = [] + n_quantizers = int(n_quantizers or self.num_quantizers) + + for quantizer in self.quantizers[:n_quantizers]: + z_q_i, indices_i, _ = quantizer(residual * mask) + quantized_out = quantized_out + z_q_i * mask + residual = residual - z_q_i * mask + all_indices.append(indices_i.astype(mx.int32)) + + if all_indices: + audio_codes = mx.stack(all_indices, axis=0) + else: + audio_codes = mx.zeros((0, batch_size, max_time), dtype=mx.int32) + + if self.output_proj is not None: + quantized_out = self.output_proj(quantized_out) + return ( + quantized_out.astype(mx.float32), + audio_codes, + input_length.astype(mx.int32), + ) + + def decode_codes(self, codes: mx.array) -> mx.array: + nq, batch_size, time_steps = codes.shape + embeddings = mx.zeros((batch_size, self.rvq_dim, time_steps), dtype=mx.float32) + for i, quantizer in enumerate(self.quantizers[:nq]): + embeddings = embeddings + quantizer.decode_code(codes[i]).astype(mx.float32) + if self.output_proj is not None: + embeddings = self.output_proj(embeddings) + return embeddings.astype(mx.float32) + + +def build_moss_audio_tokenizer_quantizer( + config: MossAudioTokenizerQuantizerConfig, +) -> MossAudioTokenizerResidualLFQ | MossAudioTokenizerResidualVQ: + quantizer_type = config.quantizer_type + if quantizer_type in {"rlfq", "random_prefix_rlfq"}: + return MossAudioTokenizerResidualLFQ( + input_dim=config.input_dim, + rvq_dim=config.rvq_dim, + output_dim=config.output_dim, + num_quantizers=config.num_quantizers, + codebook_size=config.codebook_size, + codebook_dim=config.codebook_dim, + ) + if quantizer_type in {"rvq", "spec_rvq"}: + return MossAudioTokenizerResidualVQ( + input_dim=config.input_dim, + rvq_dim=config.rvq_dim, + output_dim=config.output_dim, + num_quantizers=config.num_quantizers, + codebook_size=config.codebook_size, + codebook_dim=config.codebook_dim, + ) + raise ValueError(f"Unsupported quantizer_type: {quantizer_type}") diff --git a/mlx_audio/codec/tests/test_moss_audio_tokenizer.py b/mlx_audio/codec/tests/test_moss_audio_tokenizer.py new file mode 100644 index 000000000..efe75f83b --- /dev/null +++ b/mlx_audio/codec/tests/test_moss_audio_tokenizer.py @@ -0,0 +1,602 @@ +import json +import tempfile +import unittest +from dataclasses import asdict +from pathlib import Path +from unittest.mock import patch + +import mlx.core as mx +import numpy as np +from mlx.utils import tree_flatten + +from mlx_audio.codec.models.moss_audio_tokenizer import ( + MossAudioTokenizer, + MossAudioTokenizerConfig, + MossAudioTokenizerDecoderOutput, + MossAudioTokenizerModuleConfig, + MossAudioTokenizerQuantizerConfig, +) + + +def _tiny_moss_config() -> MossAudioTokenizerConfig: + return MossAudioTokenizerConfig( + model_type="moss_audio_tokenizer", + sampling_rate=24, + downsample_rate=4, + causal_transformer_context_duration=1.0, + encoder_modules=[ + MossAudioTokenizerModuleConfig( + module_type="PatchedPretransform", patch_size=2 + ), + MossAudioTokenizerModuleConfig( + module_type="Transformer", + input_dimension=2, + output_dimension=2, + d_model=4, + num_heads=1, + num_layers=1, + dim_feedforward=8, + causal=True, + norm="layer_norm", + positional_embedding="rope", + max_period=10000, + gating="none", + layer_scale=0.01, + conv_layout=True, + ), + MossAudioTokenizerModuleConfig( + module_type="PatchedPretransform", patch_size=2 + ), + MossAudioTokenizerModuleConfig( + module_type="Transformer", + input_dimension=4, + output_dimension=4, + d_model=4, + num_heads=1, + num_layers=1, + dim_feedforward=8, + causal=True, + norm="layer_norm", + positional_embedding="rope", + max_period=10000, + gating="none", + layer_scale=0.01, + conv_layout=True, + ), + ], + decoder_modules=[ + MossAudioTokenizerModuleConfig( + module_type="Transformer", + input_dimension=4, + output_dimension=4, + d_model=4, + num_heads=1, + num_layers=1, + dim_feedforward=8, + causal=True, + norm="layer_norm", + positional_embedding="rope", + max_period=10000, + gating="none", + layer_scale=0.01, + conv_layout=True, + ), + MossAudioTokenizerModuleConfig( + module_type="PatchedPretransform", patch_size=2 + ), + MossAudioTokenizerModuleConfig( + module_type="Transformer", + input_dimension=2, + output_dimension=2, + d_model=4, + num_heads=1, + num_layers=1, + dim_feedforward=8, + causal=True, + norm="layer_norm", + positional_embedding="rope", + max_period=10000, + gating="none", + layer_scale=0.01, + conv_layout=True, + ), + MossAudioTokenizerModuleConfig( + module_type="PatchedPretransform", patch_size=2 + ), + ], + quantizer=MossAudioTokenizerQuantizerConfig( + input_dim=4, + rvq_dim=2, + output_dim=4, + num_quantizers=2, + codebook_size=8, + codebook_dim=2, + quantizer_type="rlfq", + ), + ) + + +class TestMossAudioTokenizerModel(unittest.TestCase): + def test_encode_decode_shape_contract(self): + model = MossAudioTokenizer(_tiny_moss_config()) + audio = mx.random.normal((1, 1, 37)) + + enc = model.encode(audio, return_dict=True) + self.assertIsNotNone(enc.audio_codes) + self.assertIsNotNone(enc.audio_codes_lengths) + + assert enc.audio_codes is not None + assert enc.audio_codes_lengths is not None + self.assertEqual(enc.audio_codes.shape[0], 2) + self.assertEqual(enc.audio_codes.shape[1], 1) + self.assertEqual(int(enc.audio_codes_lengths[0]), 37 // 4) + + dec = model.decode(enc.audio_codes, return_dict=True) + self.assertIsNotNone(dec.audio) + self.assertIsNotNone(dec.audio_lengths) + + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 1) + self.assertEqual(dec.audio.shape[1], 1) + self.assertEqual(int(dec.audio_lengths[0]), int(enc.audio_codes.shape[-1]) * 4) + + def test_batch_encode_and_batch_decode(self): + model = MossAudioTokenizer(_tiny_moss_config()) + wav_a = mx.random.normal((13,)) + wav_b = mx.random.normal((27,)) + + enc = model.batch_encode([wav_a, wav_b], num_quantizers=2) + self.assertIsNotNone(enc.audio_codes) + self.assertIsNotNone(enc.audio_codes_lengths) + + assert enc.audio_codes is not None + assert enc.audio_codes_lengths is not None + self.assertEqual(enc.audio_codes.shape[:2], (2, 2)) + self.assertEqual(int(enc.audio_codes_lengths[0]), 13 // 4) + self.assertEqual(int(enc.audio_codes_lengths[1]), 27 // 4) + + codes_list = [ + enc.audio_codes[:, i, : int(enc.audio_codes_lengths[i])] + for i in range(enc.audio_codes.shape[1]) + ] + dec = model.batch_decode(codes_list, num_quantizers=2) + self.assertIsNotNone(dec.audio) + self.assertIsNotNone(dec.audio_lengths) + + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 2) + self.assertEqual(dec.audio.shape[1], 1) + self.assertEqual(int(dec.audio_lengths[0]), (13 // 4) * 4) + self.assertEqual(int(dec.audio_lengths[1]), (27 // 4) * 4) + + def test_decode_accepts_b_t_nq_layout(self): + model = MossAudioTokenizer(_tiny_moss_config()) + audio = mx.random.normal((1, 1, 40)) + enc = model.encode(audio, return_dict=True) + assert enc.audio_codes is not None + assert enc.audio_codes_lengths is not None + + codes_b_t_nq = enc.audio_codes.transpose(1, 2, 0) + dec = model.decode(codes_b_t_nq, return_dict=True) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(int(dec.audio_lengths[0]), int(enc.audio_codes_lengths[0]) * 4) + + def test_decode_accepts_quantizer_prefix_for_full_and_prefix_inputs(self): + model = MossAudioTokenizer(_tiny_moss_config()) + audio = mx.random.normal((1, 1, 40)) + enc = model.encode(audio, return_dict=True) + assert enc.audio_codes is not None + assert enc.audio_codes_lengths is not None + + dec_from_full = model.decode( + enc.audio_codes, + num_quantizers=1, + return_dict=True, + ) + dec_from_prefix = model.decode( + enc.audio_codes[:1], + num_quantizers=1, + return_dict=True, + ) + assert dec_from_full.audio is not None + assert dec_from_full.audio_lengths is not None + assert dec_from_prefix.audio is not None + assert dec_from_prefix.audio_lengths is not None + self.assertEqual( + int(dec_from_full.audio_lengths[0]), + int(dec_from_prefix.audio_lengths[0]), + ) + + def test_batch_decode_accepts_quantizer_prefix(self): + model = MossAudioTokenizer(_tiny_moss_config()) + wav_a = mx.random.normal((17,)) + wav_b = mx.random.normal((29,)) + + enc = model.batch_encode([wav_a, wav_b], num_quantizers=2) + assert enc.audio_codes is not None + assert enc.audio_codes_lengths is not None + prefix_codes_list = [ + enc.audio_codes[:1, i, : int(enc.audio_codes_lengths[i])] + for i in range(enc.audio_codes.shape[1]) + ] + dec = model.batch_decode(prefix_codes_list, num_quantizers=1) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(int(dec.audio_lengths[0]), (17 // 4) * 4) + self.assertEqual(int(dec.audio_lengths[1]), (29 // 4) * 4) + + def test_streaming_decode_matches_non_streaming_length(self): + model = MossAudioTokenizer(_tiny_moss_config()) + audio = mx.random.normal((1, 1, 64)) + enc = model.encode(audio, return_dict=True) + assert enc.audio_codes is not None + + full = model.decode(enc.audio_codes, return_dict=True) + assert full.audio is not None + assert full.audio_lengths is not None + + chunks = list(model.streaming_decode(enc.audio_codes, chunk_tokens=2)) + stream_concat = mx.concatenate(chunks, axis=-1) + self.assertEqual(stream_concat.shape[0], 1) + self.assertEqual(stream_concat.shape[1], 1) + self.assertEqual(stream_concat.shape[-1], int(full.audio_lengths[0])) + + def test_streaming_decode_accepts_quantizer_prefix(self): + model = MossAudioTokenizer(_tiny_moss_config()) + audio = mx.random.normal((1, 1, 64)) + enc = model.encode(audio, return_dict=True) + assert enc.audio_codes is not None + + full = model.decode(enc.audio_codes, num_quantizers=1, return_dict=True) + assert full.audio is not None + assert full.audio_lengths is not None + + chunks = list( + model.streaming_decode( + enc.audio_codes[:1], chunk_tokens=2, num_quantizers=1 + ) + ) + stream_concat = mx.concatenate(chunks, axis=-1) + self.assertEqual(stream_concat.shape[-1], int(full.audio_lengths[0])) + + def test_decode_accepts_encode_output_when_time_equals_quantizers(self): + model = MossAudioTokenizer(_tiny_moss_config()) + # downsample_rate=4 => 8 samples encodes to 2 frames, matching num_quantizers=2. + audio = mx.random.normal((1, 1, 8)) + enc = model.encode(audio, return_dict=True) + assert enc.audio_codes is not None + assert enc.audio_codes_lengths is not None + self.assertEqual(tuple(enc.audio_codes.shape), (2, 1, 2)) + + dec = model.decode(enc.audio_codes, return_dict=True) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 1) + self.assertEqual(int(dec.audio_lengths[0]), int(enc.audio_codes_lengths[0]) * 4) + + def test_decode_prefers_nq_first_when_3d_shape_matches_both_orientations(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((2, 5, 2), dtype=mx.int32) + + dec = model.decode(tie_shape_codes, return_dict=True) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 5) + self.assertTrue(np.all(np.array(dec.audio_lengths) == 8)) + + def test_decode_preserves_canonical_3d_tie_for_explicit_configured_num_quantizers( + self, + ): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.array( + [ + [[1, 2], [3, 4], [5, 6]], + [[7, 8], [9, 10], [11, 12]], + ], + dtype=mx.int32, + ) + captured: list[np.ndarray] = [] + + def fake_decode_frame(audio_codes, audio_codes_lengths=None, caches=None): + del audio_codes_lengths, caches + captured.append(np.array(audio_codes)) + batch = int(audio_codes.shape[1]) + time_steps = int(audio_codes.shape[2]) + sample_count = time_steps * model.downsample_rate + return MossAudioTokenizerDecoderOutput( + audio=mx.zeros((batch, 1, sample_count), dtype=mx.float32), + audio_lengths=mx.full((batch,), sample_count, dtype=mx.int32), + ) + + with patch.object(model, "_decode_frame", side_effect=fake_decode_frame): + dec = model.decode(tie_shape_codes, num_quantizers=2, return_dict=True) + + assert dec.audio is not None + self.assertEqual(len(captured), 1) + np.testing.assert_array_equal(captured[0], np.array(tie_shape_codes)) + + def test_decode_prefers_requested_quantizer_match_for_2d_nq_last_prefix_tie(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((2, 1), dtype=mx.int32) + + dec = model.decode(tie_shape_codes, num_quantizers=1, return_dict=True) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 1) + self.assertEqual(int(dec.audio_lengths[0]), 8) + + def test_decode_prefers_requested_quantizer_match_for_3d_nq_last_prefix_tie(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((2, 1, 1), dtype=mx.int32) + + dec = model.decode(tie_shape_codes, num_quantizers=1, return_dict=True) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 2) + self.assertTrue(np.all(np.array(dec.audio_lengths) == 4)) + + def test_decode_preserves_canonical_orientation_for_true_prefix_3d_tie(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.array( + [ + [[1], [2], [3]], + ], + dtype=mx.int32, + ) + captured: list[np.ndarray] = [] + + def fake_decode_frame(audio_codes, audio_codes_lengths=None, caches=None): + del audio_codes_lengths, caches + captured.append(np.array(audio_codes)) + batch = int(audio_codes.shape[1]) + time_steps = int(audio_codes.shape[2]) + sample_count = time_steps * model.downsample_rate + return MossAudioTokenizerDecoderOutput( + audio=mx.zeros((batch, 1, sample_count), dtype=mx.float32), + audio_lengths=mx.full((batch,), sample_count, dtype=mx.int32), + ) + + with patch.object(model, "_decode_frame", side_effect=fake_decode_frame): + dec = model.decode(tie_shape_codes, num_quantizers=1, return_dict=True) + + assert dec.audio is not None + self.assertEqual(len(captured), 1) + np.testing.assert_array_equal(captured[0], np.array(tie_shape_codes)) + + def test_decode_preserves_nq_first_for_explicit_square_tie(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.array([[1, 2], [3, 4]], dtype=mx.int32) + captured: list[np.ndarray] = [] + + def fake_decode_frame(audio_codes, audio_codes_lengths=None, caches=None): + del audio_codes_lengths, caches + captured.append(np.array(audio_codes)) + batch = int(audio_codes.shape[1]) + time_steps = int(audio_codes.shape[2]) + sample_count = time_steps * model.downsample_rate + return MossAudioTokenizerDecoderOutput( + audio=mx.zeros((batch, 1, sample_count), dtype=mx.float32), + audio_lengths=mx.full((batch,), sample_count, dtype=mx.int32), + ) + + with patch.object(model, "_decode_frame", side_effect=fake_decode_frame): + dec = model.decode(tie_shape_codes, num_quantizers=2, return_dict=True) + + assert dec.audio is not None + self.assertEqual(len(captured), 1) + expected = np.array([[[1, 2]], [[3, 4]]], dtype=np.int32) + np.testing.assert_array_equal(captured[0], expected) + + def test_batch_decode_prefers_requested_quantizer_match_for_2d_nq_last_prefix_tie( + self, + ): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((2, 1), dtype=mx.int32) + + dec = model.batch_decode([tie_shape_codes], num_quantizers=1) + assert dec.audio is not None + assert dec.audio_lengths is not None + self.assertEqual(dec.audio.shape[0], 1) + self.assertEqual(int(dec.audio_lengths[0]), 8) + + def test_batch_decode_rejects_canonical_batched_true_prefix_3d_tie( + self, + ): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((1, 3, 1), dtype=mx.int32) + + with self.assertRaisesRegex( + ValueError, + "batch_decode\\(\\) expects each codes_list entry to resolve to batch_size=1", + ): + _ = model.batch_decode([tie_shape_codes], num_quantizers=1) + + def test_batch_decode_preserves_nq_first_for_explicit_square_tie(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.array([[1, 2], [3, 4]], dtype=mx.int32) + captured: list[np.ndarray] = [] + + def fake_decode_frame(audio_codes, audio_codes_lengths=None, caches=None): + del audio_codes_lengths, caches + captured.append(np.array(audio_codes)) + batch = int(audio_codes.shape[1]) + time_steps = int(audio_codes.shape[2]) + sample_count = time_steps * model.downsample_rate + return MossAudioTokenizerDecoderOutput( + audio=mx.zeros((batch, 1, sample_count), dtype=mx.float32), + audio_lengths=mx.full((batch,), sample_count, dtype=mx.int32), + ) + + with patch.object(model, "_decode_frame", side_effect=fake_decode_frame): + dec = model.batch_decode([tie_shape_codes], num_quantizers=2) + + assert dec.audio is not None + self.assertEqual(len(captured), 1) + expected = np.array([[[1, 2]], [[3, 4]]], dtype=np.int32) + np.testing.assert_array_equal(captured[0], expected) + + def test_streaming_decode_prefers_requested_quantizer_match_for_2d_nq_last_prefix_tie( + self, + ): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((2, 1), dtype=mx.int32) + + chunks = list( + model.streaming_decode( + tie_shape_codes, + chunk_tokens=1, + num_quantizers=1, + ) + ) + stream_concat = mx.concatenate(chunks, axis=-1) + self.assertEqual(stream_concat.shape[0], 1) + self.assertEqual(stream_concat.shape[1], 1) + self.assertEqual(stream_concat.shape[-1], 8) + + def test_streaming_decode_preserves_nq_first_for_explicit_square_tie(self): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.array( + [ + [[1, 2]], + [[3, 4]], + ], + dtype=mx.int32, + ) + captured: list[np.ndarray] = [] + + def fake_decode_frame(audio_codes, audio_codes_lengths=None, caches=None): + del audio_codes_lengths, caches + captured.append(np.array(audio_codes)) + batch = int(audio_codes.shape[1]) + time_steps = int(audio_codes.shape[2]) + sample_count = time_steps * model.downsample_rate + return MossAudioTokenizerDecoderOutput( + audio=mx.zeros((batch, 1, sample_count), dtype=mx.float32), + audio_lengths=mx.full((batch,), sample_count, dtype=mx.int32), + ) + + with patch.object(model, "_decode_frame", side_effect=fake_decode_frame): + _ = list( + model.streaming_decode( + tie_shape_codes, + chunk_tokens=1, + num_quantizers=2, + ) + ) + + self.assertGreaterEqual(len(captured), 1) + reconstructed = np.concatenate(captured, axis=2) + expected = np.array([[[1, 2]], [[3, 4]]], dtype=np.int32) + np.testing.assert_array_equal(reconstructed, expected) + + def test_streaming_decode_rejects_canonical_batched_true_prefix_3d_tie( + self, + ): + model = MossAudioTokenizer(_tiny_moss_config()) + tie_shape_codes = mx.zeros((1, 3, 1), dtype=mx.int32) + + with self.assertRaisesRegex( + ValueError, + "streaming_decode currently only supports batch_size=1", + ): + _ = list( + model.streaming_decode( + tie_shape_codes, + chunk_tokens=1, + num_quantizers=1, + ) + ) + + def test_sanitize_reconstructs_weight_norm(self): + model = MossAudioTokenizer(_tiny_moss_config()) + expected_shapes = { + name: tuple(value.shape) for name, value in tree_flatten(model.parameters()) + } + linear_key = next( + key + for key in expected_shapes + if key.endswith("encoder.1.transformer.layers.0.linear1.weight") + ) + linear_shape = expected_shapes[linear_key] + + expected_input_proj_shape = expected_shapes["quantizer.input_proj.weight"] + pytorch_v_shape = ( + expected_input_proj_shape[0], + expected_input_proj_shape[2], + expected_input_proj_shape[1], + ) + g = mx.ones((pytorch_v_shape[0], 1, 1), dtype=mx.float32) + v = mx.arange( + 1, + 1 + int(np.prod(np.array(pytorch_v_shape))), + dtype=mx.float32, + ).reshape(pytorch_v_shape) + weights = { + "quantizer.input_proj.parametrizations.weight.original0": g, + "quantizer.input_proj.parametrizations.weight.original1": v, + linear_key: mx.ones(linear_shape, dtype=mx.float32), + } + sanitized = model.sanitize(weights) + + self.assertIn("quantizer.input_proj.weight", sanitized) + self.assertEqual( + sanitized["quantizer.input_proj.weight"].shape, + expected_shapes["quantizer.input_proj.weight"], + ) + self.assertNotIn( + "quantizer.input_proj.parametrizations.weight.original0", sanitized + ) + self.assertNotIn( + "quantizer.input_proj.parametrizations.weight.original1", sanitized + ) + self.assertTrue( + np.allclose(np.array(sanitized[linear_key]), np.array(weights[linear_key])) + ) + + def test_from_pretrained_loads_local_directory(self): + model = MossAudioTokenizer(_tiny_moss_config()) + with tempfile.TemporaryDirectory() as tmp_dir: + root = Path(tmp_dir) + config_path = root / "config.json" + config_payload = asdict(model.config) + config_payload["model_type"] = "speech_tokenizer" + config_path.write_text(json.dumps(config_payload), encoding="utf-8") + + weight_path = root / "model.safetensors" + mx.save_safetensors( + weight_path.as_posix(), + dict(tree_flatten(model.parameters())), + ) + + loaded = MossAudioTokenizer.from_pretrained(root) + audio = mx.random.normal((1, 1, 20)) + enc = loaded.encode(audio, return_dict=True) + self.assertIsNotNone(enc.audio_codes) + self.assertIsNotNone(enc.audio_codes_lengths) + + +class TestMossAudioTokenizerQuantPredicate(unittest.TestCase): + def test_model_quant_predicate_skips_embeddings(self): + model = MossAudioTokenizer(_tiny_moss_config()) + codebook_module = model.quantizer.quantizers[0].codebook + linear_module = model.encoder[1].transformer.layers[0].linear1 + + self.assertFalse( + model.model_quant_predicate( + "quantizer.quantizers.0.codebook", + codebook_module, + ) + ) + self.assertTrue( + model.model_quant_predicate( + "encoder.1.transformer.layers.0.linear1", + linear_module, + ) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/mlx_audio/codec/tests/test_moss_audio_tokenizer_config_contracts.py b/mlx_audio/codec/tests/test_moss_audio_tokenizer_config_contracts.py new file mode 100644 index 000000000..d1cbcc48b --- /dev/null +++ b/mlx_audio/codec/tests/test_moss_audio_tokenizer_config_contracts.py @@ -0,0 +1,45 @@ +import unittest +from pathlib import Path + +from mlx_audio.codec.models.moss_audio_tokenizer.config import ( + CANONICAL_MODEL_TYPE, + load_moss_audio_tokenizer_config, +) + + +def _repo_root() -> Path: + # .../mlx_audio/codec/tests -> project root + return Path(__file__).resolve().parents[3] + + +def _fixture_path() -> Path: + return ( + _repo_root() + / "tests" + / "fixtures" + / "moss" + / "MOSS-Audio-Tokenizer" + / "config.json" + ) + + +class TestMossAudioTokenizerPhase0Config(unittest.TestCase): + def test_load_upstream_config_contract(self): + config = load_moss_audio_tokenizer_config(_fixture_path()) + self.assertEqual(config.model_type, CANONICAL_MODEL_TYPE) + self.assertEqual(config.sampling_rate, 24000) + self.assertEqual(config.downsample_rate, 1920) + self.assertEqual(config.frame_rate, 12.5) + self.assertEqual(config.quantizer.num_quantizers, 32) + self.assertEqual(config.quantizer.codebook_size, 1024) + self.assertEqual(config.quantizer.quantizer_type, "rlfq") + + def test_patch_products_match_downsample_rate(self): + config = load_moss_audio_tokenizer_config(_fixture_path()) + self.assertEqual(config.encoder_patch_product, 1920) + self.assertEqual(config.decoder_patch_product, 1920) + self.assertTrue(config.patch_alignment_is_valid()) + + +if __name__ == "__main__": + unittest.main() diff --git a/mlx_audio/convert.py b/mlx_audio/convert.py index 256a25de5..6aa389136 100644 --- a/mlx_audio/convert.py +++ b/mlx_audio/convert.py @@ -152,9 +152,14 @@ def _discover_detection_hints(domain: str) -> dict: module_path = f"mlx_audio.{domain}.models.{model_type}" try: module = importlib.import_module(module_path) + has_model_entrypoint = hasattr(module, "Model") + has_model_config = hasattr(module, "ModelConfig") # Check for explicit detection hints if hasattr(module, "DETECTION_HINTS"): + if not has_model_entrypoint: + # Bootstrap-only modules should not participate in autodetection. + continue model_hints = module.DETECTION_HINTS if "config_keys" in model_hints: hints["config_keys"][model_type] = set(model_hints["config_keys"]) @@ -168,15 +173,16 @@ def _discover_detection_hints(domain: str) -> dict: ) else: # Infer from ModelConfig if available - if hasattr(module, "ModelConfig"): + if has_model_config: config_keys = _get_config_keys(module.ModelConfig) hints["config_keys"][model_type] = config_keys # Use model_type as default path pattern - hints["path_patterns"][model_type] = { - model_type, - model_type.replace("_", ""), - } + if has_model_entrypoint: + hints["path_patterns"][model_type] = { + model_type, + model_type.replace("_", ""), + } except ImportError: continue @@ -565,6 +571,11 @@ def convert( # Get model class and instantiate model_class = get_model_class(model_type, domain) + if not hasattr(model_class, "Model"): + raise ValueError( + f"Model module '{domain.value}.models.{model_type}' does not expose " + "a Model entry point yet." + ) model_config = ( model_class.ModelConfig.from_dict(config) diff --git a/mlx_audio/tts/generate.py b/mlx_audio/tts/generate.py index de3e03a4a..36080e97f 100644 --- a/mlx_audio/tts/generate.py +++ b/mlx_audio/tts/generate.py @@ -1,7 +1,8 @@ import argparse +import json import os import sys -from typing import Optional, Tuple, Union +from typing import Any, Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn @@ -14,6 +15,177 @@ from .audio_player import AudioPlayer from .utils import load_model +_MOSS_PRESET_ALIASES = { + "voicegenerator": "voice_generator", + "moss_voice_generator": "voice_generator", + "moss_voicegenerator": "voice_generator", + "moss_voice_design": "voice_generator", + "moss_sound_effect": "soundeffect", + "sound_effect": "soundeffect", + "moss_soundeffect": "soundeffect", +} + +_MOSS_PRESET_TASKS = { + "moss_tts": "default_tts", + "moss_tts_local": "default_tts", + "ttsd": "ttsd", + "soundeffect": "soundeffect", + "voice_generator": "voice_generator", + "realtime": "realtime", +} + +_TASK_DEFAULT_PRESETS = { + "ttsd": "ttsd", + "soundeffect": "soundeffect", + "voice_generator": "voice_generator", + "realtime": "realtime", +} + +_REALTIME_ONLY_KWARGS = frozenset( + { + "chunk_frames", + "overlap_frames", + "decode_chunk_duration", + "max_pending_frames", + "repetition_window", + } +) + + +def _normalize_preset(preset: Optional[str]) -> Optional[str]: + if preset is None: + return None + normalized = str(preset).strip().lower().replace("-", "_") + return _MOSS_PRESET_ALIASES.get(normalized, normalized) + + +def _normalize_model_type(model_type: Optional[str]) -> Optional[str]: + if model_type is None: + return None + return str(model_type).strip().lower().replace("-", "_") + + +def _looks_like_moss_model(model: Optional[Union[str, nn.Module]]) -> bool: + if isinstance(model, str): + return "moss" in model.lower() + if model is None: + return False + model_type = _normalize_model_type(getattr(model, "model_type", None)) + if model_type: + return "moss" in model_type + return "moss" in str(type(model)).lower() + + +def _is_explicit_realtime_model(model: Optional[Union[str, nn.Module]]) -> bool: + if isinstance(model, str): + normalized = model.strip().lower().replace("-", "_") + return "moss_tts_realtime" in normalized or "moss/tts/realtime" in normalized + if model is None: + return False + model_type = _normalize_model_type(getattr(model, "model_type", None)) + return model_type == "moss_tts_realtime" + + +def _parse_model_kwargs_json( + model_kwargs_json: Optional[Union[str, dict[str, Any]]], +) -> dict[str, Any]: + if model_kwargs_json is None: + return {} + parsed_kwargs = model_kwargs_json + if isinstance(parsed_kwargs, str): + try: + parsed_kwargs = json.loads(parsed_kwargs) + except json.JSONDecodeError as exc: + raise ValueError(f"Invalid --model_kwargs_json value: {exc.msg}") from exc + if not isinstance(parsed_kwargs, dict): + raise ValueError("--model_kwargs_json must decode to a JSON object") + return dict(parsed_kwargs) + + +def _load_dialogue_speakers( + dialogue_speakers_json: Optional[str], +) -> Optional[list[dict[str, Any]]]: + if dialogue_speakers_json is None: + return None + with open(dialogue_speakers_json, "r", encoding="utf-8") as f: + raw = json.load(f) + if not isinstance(raw, list): + raise ValueError("dialogue_speakers_json must contain a JSON list") + dialogue_speakers: list[dict[str, Any]] = [] + for item in raw: + if not isinstance(item, dict): + raise ValueError("Each dialogue_speakers item must be a JSON object") + dialogue_speakers.append(item) + return dialogue_speakers + + +def _infer_gateway_task_and_preset( + *, + model: Optional[Union[str, nn.Module]], + preset: Optional[str], + dialogue_speakers: Optional[list[dict[str, Any]]], + sound_event: Optional[str], + ambient_sound: Optional[str], + instruct: Optional[str], + repetition_window: Optional[int], + passthrough_kwargs: dict[str, Any], +) -> tuple[str, Optional[str]]: + normalized_preset = _normalize_preset(preset) + explicit_preset_task = _MOSS_PRESET_TASKS.get(normalized_preset) + + explicit_realtime_marker = _is_explicit_realtime_model(model) + realtime_kwarg_marker = bool(_REALTIME_ONLY_KWARGS & set(passthrough_kwargs.keys())) + moss_hint = ( + _looks_like_moss_model(model) + or normalized_preset in _MOSS_PRESET_TASKS + or bool(dialogue_speakers) + or bool(ambient_sound or sound_event) + ) + realtime_marker = ( + explicit_realtime_marker + or normalized_preset == "realtime" + or (moss_hint and (repetition_window is not None or realtime_kwarg_marker)) + ) + + ttsd_marker = bool(dialogue_speakers) or normalized_preset == "ttsd" + soundeffect_marker = bool(ambient_sound or sound_event) or ( + normalized_preset == "soundeffect" + ) + + moss_context = moss_hint or ttsd_marker or soundeffect_marker or realtime_marker + has_instruct = bool(str(instruct).strip()) if instruct is not None else False + voice_design_marker = normalized_preset == "voice_generator" or ( + moss_context and has_instruct + ) + + inferred_task = "default_tts" + if realtime_marker: + inferred_task = "realtime" + elif ttsd_marker: + inferred_task = "ttsd" + elif soundeffect_marker: + inferred_task = "soundeffect" + elif voice_design_marker: + inferred_task = "voice_generator" + + if ( + explicit_preset_task is not None + and inferred_task != explicit_preset_task + and inferred_task != "default_tts" + ): + raise ValueError( + "Incompatible task markers: " + f"preset={preset!r} conflicts with inferred task '{inferred_task}'" + ) + + effective_preset = preset + if normalized_preset in _MOSS_PRESET_TASKS: + effective_preset = normalized_preset + elif inferred_task in _TASK_DEFAULT_PRESETS: + effective_preset = _TASK_DEFAULT_PRESETS[inferred_task] + + return inferred_task, effective_preset + def detect_speech_boundaries( wav: np.ndarray, @@ -103,11 +275,23 @@ def hertz_to_mel(pitch: float) -> float: def generate_audio( - text: str, + text: Optional[str], model: Optional[Union[str, nn.Module]] = None, max_tokens: int = 1200, + tokens: Optional[int] = None, + duration_s: Optional[float] = None, + seconds: Optional[float] = None, + n_vq_for_inference: Optional[int] = None, voice: str = "af_heart", instruct: Optional[str] = None, + quality: Optional[str] = None, + sound_event: Optional[str] = None, + ambient_sound: Optional[str] = None, + language: Optional[str] = None, + preset: Optional[str] = None, + model_kwargs_json: Optional[Union[str, dict[str, Any]]] = None, + dialogue_speakers_json: Optional[str] = None, + input_type: str = "text", speed: float = 1.0, lang_code: str = "en", cfg_scale: Optional[float] = None, @@ -124,8 +308,18 @@ def generate_audio( play: bool = False, verbose: bool = True, temperature: float = 0.7, + seed: Optional[int] = None, + repetition_window: Optional[int] = None, stream: bool = False, streaming_interval: float = 2.0, + long_form: bool = False, + long_form_min_chars: int = 160, + long_form_target_chars: int = 320, + long_form_max_chars: int = 520, + long_form_prefix_audio_seconds: float = 2.0, + long_form_prefix_audio_max_tokens: int = 25, + long_form_prefix_text_chars: int = 0, + long_form_retry_attempts: int = 0, **kwargs, ) -> None: """ @@ -154,7 +348,15 @@ def generate_audio( - None: The function writes the generated audio to a file. """ try: - play = play or stream + if seed is not None: + mx.random.seed(int(seed)) + + # Keep streaming generation usable in headless/non-interactive runs. + # Playback remains explicitly opt-in via --play. + play = bool(play) + + if (text is None or text.strip() == "") and ambient_sound: + text = ambient_sound if model is None: raise ValueError("Model path or model instance must be provided.") @@ -205,6 +407,24 @@ def generate_audio( if output_path: os.makedirs(output_path, exist_ok=True) file_prefix = os.path.join(output_path, file_prefix) + has_stream_output_sink = bool(output_path) + if stream and not play and not has_stream_output_sink: + raise ValueError( + "Streaming mode requires at least one sink: enable --play or provide --output_path." + ) + + parsed_model_kwargs_json = _parse_model_kwargs_json(model_kwargs_json) + dialogue_speakers = _load_dialogue_speakers(dialogue_speakers_json) + _, effective_preset = _infer_gateway_task_and_preset( + model=model, + preset=preset, + dialogue_speakers=dialogue_speakers, + sound_event=sound_event, + ambient_sound=ambient_sound, + instruct=instruct, + repetition_window=repetition_window, + passthrough_kwargs=kwargs, + ) if instruct is not None: print(f"\033[94mInstruct:\033[0m {instruct}") @@ -225,7 +445,18 @@ def generate_audio( ref_text=ref_text, cfg_scale=cfg_scale, ddpm_steps=ddpm_steps, + tokens=tokens, + duration_s=duration_s if duration_s is not None else seconds, + quality=quality, + sound_event=sound_event, + ambient_sound=ambient_sound, + language=language, + preset=effective_preset, + n_vq_for_inference=n_vq_for_inference, + dialogue_speakers=dialogue_speakers, + input_type=input_type, temperature=temperature, + repetition_window=repetition_window, max_tokens=max_tokens, verbose=verbose, stream=stream, @@ -233,6 +464,25 @@ def generate_audio( instruct=instruct, **kwargs, ) + gen_kwargs.update(parsed_model_kwargs_json) + + if long_form: + gen_kwargs.update( + { + "long_form": True, + "long_form_min_chars": int(long_form_min_chars), + "long_form_target_chars": int(long_form_target_chars), + "long_form_max_chars": int(long_form_max_chars), + "long_form_prefix_audio_seconds": float( + long_form_prefix_audio_seconds + ), + "long_form_prefix_audio_max_tokens": int( + long_form_prefix_audio_max_tokens + ), + "long_form_prefix_text_chars": int(long_form_prefix_text_chars), + "long_form_retry_attempts": int(long_form_retry_attempts), + } + ) results = model.generate(**gen_kwargs) @@ -242,9 +492,20 @@ def generate_audio( if play: player.queue_audio(result.audio) - if join_audio: + if join_audio and not stream: audio_list.append(result.audio) - elif not stream: + if stream and has_stream_output_sink: + file_name = f"{file_prefix}_{i:03d}.{audio_format}" + audio_write( + file_name, + np.array(result.audio), + result.sample_rate, + format=audio_format, + ) + print( + f"✅ Stream chunk successfully generated and saved as: {file_name}" + ) + elif not stream and not join_audio: file_name = f"{file_prefix}_{i:03d}.{audio_format}" audio_write( file_name, @@ -274,14 +535,18 @@ def generate_audio( if join_audio and not stream: if verbose: print(f"Joining {len(audio_list)} audio files") + joined_file_name = f"{file_prefix}.{audio_format}" audio = mx.concatenate(audio_list, axis=0) audio_write( - f"{file_prefix}.{audio_format}", + joined_file_name, audio, model.sample_rate, + format=audio_format, ) if verbose: - print(f"✅ Audio successfully generated and saving as: {file_name}") + print( + f"✅ Audio successfully generated and saving as: {joined_file_name}" + ) if play: player.wait_for_drain() @@ -299,6 +564,11 @@ def generate_audio( traceback.print_exc() +def generate_stream(text: Optional[str], **kwargs) -> None: + kwargs["stream"] = True + generate_audio(text=text, **kwargs) + + def parse_args(): parser = argparse.ArgumentParser(description="Generate audio from text using TTS.") parser.add_argument( @@ -313,6 +583,32 @@ def parse_args(): default=1200, help="Maximum number of tokens to generate", ) + parser.add_argument( + "--tokens", + type=int, + default=None, + help="Target duration control for models that support token-based timing", + ) + parser.add_argument( + "--duration_s", + "--seconds", + dest="duration_s", + type=float, + default=None, + help=( + "Convenience duration control in seconds (mapped to tokens at 12.5 Hz; " + "ignored when --tokens is provided)" + ), + ) + parser.add_argument( + "--n_vq_for_inference", + type=int, + default=None, + help=( + "Local-only inference depth override (1..n_vq) for quality/performance " + "trade-offs" + ), + ) parser.add_argument( "--text", type=str, @@ -331,6 +627,51 @@ def parse_args(): default=None, help="Instruction for CustomVoice (emotion/style) or VoiceDesign (voice description)", ) + parser.add_argument( + "--quality", type=str, default=None, help="Quality hint for supported models" + ) + parser.add_argument( + "--sound_event", + type=str, + default=None, + help="Sound event description for supported models", + ) + parser.add_argument( + "--ambient_sound", + type=str, + default=None, + help="Ambient sound description for supported models", + ) + parser.add_argument( + "--language", + type=str, + default=None, + help="Language hint for supported models", + ) + parser.add_argument( + "--preset", + type=str, + default=None, + help=( + "Variant sampling preset (MOSS family): moss_tts, moss_tts_local, " + "ttsd, voice_generator, soundeffect, realtime" + ), + ) + parser.add_argument( + "--model_kwargs_json", + type=str, + default=None, + help="JSON object for advanced model.generate kwargs (escape hatch)", + ) + parser.add_argument( + "--dialogue_speakers_json", + type=str, + default=None, + help=( + "Path to TTSD speaker schema JSON " + "(list of {speaker_id, ref_audio, ref_text/text})" + ), + ) parser.add_argument( "--exaggeration", type=float, @@ -356,6 +697,13 @@ def parse_args(): ) parser.add_argument("--pitch", type=float, default=1.0, help="Pitch of the voice") parser.add_argument("--lang_code", type=str, default="en", help="Language code") + parser.add_argument( + "--input_type", + type=str, + default="text", + choices=["text", "pinyin", "ipa"], + help="Input representation for supported models", + ) parser.add_argument( "--output_path", type=str, default=None, help="Directory path for output files" ) @@ -386,6 +734,12 @@ def parse_args(): parser.add_argument( "--temperature", type=float, default=0.7, help="Temperature for the model" ) + parser.add_argument( + "--seed", + type=int, + default=None, + help="Optional random seed for reproducible sampling paths", + ) parser.add_argument("--top_p", type=float, default=0.9, help="Top-p for the model") parser.add_argument("--top_k", type=int, default=50, help="Top-k for the model") parser.add_argument( @@ -394,6 +748,15 @@ def parse_args(): default=1.1, help="Repetition penalty for the model", ) + parser.add_argument( + "--repetition_window", + type=int, + default=None, + help=( + "Realtime repetition-history window size; <=0 disables windowing " + "and applies repetition penalty over full history" + ), + ) parser.add_argument( "--stream", action="store_true", @@ -405,9 +768,59 @@ def parse_args(): default=2.0, help="The time interval in seconds for streaming segments", ) + parser.add_argument( + "--long_form", + action="store_true", + help="Enable segmented long-form generation for MOSS-TTS variants", + ) + parser.add_argument( + "--long_form_min_chars", + type=int, + default=160, + help="Minimum per-segment text budget for long-form planning", + ) + parser.add_argument( + "--long_form_target_chars", + type=int, + default=320, + help="Target per-segment text budget for long-form planning", + ) + parser.add_argument( + "--long_form_max_chars", + type=int, + default=520, + help="Maximum per-segment text budget for long-form planning", + ) + parser.add_argument( + "--long_form_prefix_audio_seconds", + type=float, + default=2.0, + help="Carry-forward tail duration (seconds) between long-form segments", + ) + parser.add_argument( + "--long_form_prefix_audio_max_tokens", + type=int, + default=25, + help="Carry-forward tail budget in audio tokens (stricter cap wins)", + ) + parser.add_argument( + "--long_form_prefix_text_chars", + type=int, + default=0, + help="Optional carry-forward text window size in characters", + ) + parser.add_argument( + "--long_form_retry_attempts", + type=int, + default=0, + help="Retry attempts per long-form segment before failing", + ) args = parser.parse_args() + if args.text is None and args.ambient_sound is not None: + args.text = args.ambient_sound + if args.text is None: if not sys.stdin.isatty(): args.text = sys.stdin.read().strip() diff --git a/mlx_audio/tts/models/moss_tts/README.md b/mlx_audio/tts/models/moss_tts/README.md new file mode 100644 index 000000000..84b31a160 --- /dev/null +++ b/mlx_audio/tts/models/moss_tts/README.md @@ -0,0 +1,320 @@ +# MOSS-TTS Family (MLX Runtime) + +Unified MLX runtime support for the OpenMOSS MOSS-TTS family. + +Supported checkpoints: + +- `OpenMOSS-Team/MOSS-TTS` (Delay) +- `OpenMOSS-Team/MOSS-TTS-Local-Transformer` (Local) +- `OpenMOSS-Team/MOSS-TTSD-v1.0` (TTSD) +- `OpenMOSS-Team/MOSS-Voice-Generator` (VoiceGenerator) +- `OpenMOSS-Team/MOSS-SoundEffect` (SoundEffect) + +Realtime has a dedicated runtime package: `mlx_audio/tts/models/moss_tts_realtime/`. + +## Model Variants + +| Variant | HF ID | Runtime | Preset | Primary Use | +|---|---|---|---|---| +| Delay | `OpenMOSS-Team/MOSS-TTS` | `moss_tts` | `moss_tts` | General high-capability TTS | +| Local | `OpenMOSS-Team/MOSS-TTS-Local-Transformer` | `moss_tts` | `moss_tts_local` | Lower-memory TTS, inference depth override | +| TTSD | `OpenMOSS-Team/MOSS-TTSD-v1.0` | `moss_tts` | `ttsd` | Multi-speaker dialogue | +| VoiceGenerator | `OpenMOSS-Team/MOSS-Voice-Generator` | `moss_tts` | `voice_generator` | Voice design from text instruction | +| SoundEffect | `OpenMOSS-Team/MOSS-SoundEffect` | `moss_tts` | `soundeffect` | Text-to-sound-event synthesis | + +## Quick Start + +### Python API + +```python +from mlx_audio.tts.utils import load_model + +model = load_model("OpenMOSS-Team/MOSS-TTS-Local-Transformer") + +results = list( + model.generate( + text="Hello from MOSS on MLX.", + preset="moss_tts_local", + input_type="text", # text | pinyin | ipa + duration_s=6.0, # mapped to tokens at 12.5 Hz when tokens is omitted + max_tokens=240, + ) +) + +audio = results[0].audio +``` + +### CLI + +```bash +uv run python -m mlx_audio.tts.generate \ + --model OpenMOSS-Team/MOSS-TTS-Local-Transformer \ + --text "Hello from MOSS on MLX." \ + --preset moss_tts_local \ + --duration_s 6 \ + --output_path ./outputs/moss_local +``` + +More complete runnable examples are in: + +- `examples/moss_tts_basic.py` +- `examples/moss_tts_deterministic.py` +- `examples/moss_tts_voice_cloning.py` +- `examples/moss_tts_continuation_showcase.py` +- `examples/moss_ttsd_dialogue.py` +- `examples/moss_voice_design.py` +- `examples/moss_sound_effects.py` +- `examples/moss_tts_long_form.py` +- `examples/moss_tts_pronunciation_control.py` +- `examples/moss_tts_realtime_text_deltas.py` +- `examples/moss_tts_realtime_multiturn_agent.py` +- `examples/moss_tts_showcase_album.py` + +## Showcase Recipes + +Use these scripts as fast entry points for the most important MOSS workflows: + +- Continuation prompting (assistant prefix audio, no `ref_audio` generate arg): + +```bash +uv run python examples/moss_tts_continuation_showcase.py \ + --model OpenMOSS-Team/MOSS-TTS-Local-Transformer \ + --preset moss_tts_local +``` + +- Realtime multiturn agent flow (persistent voice prompt + per-turn user audio): + +```bash +uv run python examples/moss_tts_realtime_multiturn_agent.py \ + --model OpenMOSS-Team/MOSS-TTS-Realtime \ + --save-chunks +``` + +- Full family album export (all variants + shareable manifest files): + +```bash +uv run python examples/moss_tts_showcase_album.py \ + --output-dir outputs/moss_tts_showcase_album +``` + +`moss_tts_showcase_album.py` writes `showcase_album.json` and +`showcase_album.md` alongside generated tracks. + +## Deterministic Recipes + +If voice/style/emotion vary between runs, use the deterministic script: +`examples/moss_tts_deterministic.py`. + +Default mode is deterministic **hybrid** decoding: + +- text/control channel: greedy (`do_sample=False`) +- audio channels: seeded sampling (`do_sample=True`) + +This remains reproducible but avoids a known full-greedy silence collapse on +some prompts/checkpoints. + +- Deterministic text-only generation: + +```bash +uv run python examples/moss_tts_deterministic.py \ + --model OpenMOSS-Team/MOSS-TTS-Local-Transformer \ + --preset moss_tts_local \ + --text "Hello what is happening this is from MOSS on MLX." \ + --instruct "Calm, friendly, medium speaking rate, conversational tone." \ + --output-dir outputs/moss_tts_deterministic +``` + +- Direct CLI equivalent (`-m mlx_audio.tts.generate`) using `model_kwargs_json`: + +```bash +uv run python -m mlx_audio.tts.generate \ + --model OpenMOSS-Team/MOSS-TTS-Local-Transformer \ + --preset moss_tts_local \ + --text "Hello what is happening this is from MOSS on MLX." \ + --instruct "Calm, friendly, medium speaking rate, conversational tone." \ + --seed 1234 \ + --model_kwargs_json '{"do_samples":[false]}' \ + --output_path outputs/moss_tts_deterministic_cli +``` + +- Deterministic reference-conditioned generation (strongest voice consistency): + +```bash +uv run python examples/moss_tts_deterministic.py \ + --model OpenMOSS-Team/MOSS-TTS-Local-Transformer \ + --preset moss_tts_local \ + --with-reference \ + --ref-audio REFERENCE/MOSS-Audio-Tokenizer/demo/demo_gt.wav \ + --ref-text "Demo reference transcript." \ + --text "Hello what is happening this is from MOSS on MLX." \ + --instruct "Warm expressive narrator, medium speaking rate." \ + --output-dir outputs/moss_tts_deterministic_ref +``` + +Notes: + +- `preset` sets default sampling (`temperature`, `top_p`, `top_k`, `repetition_penalty`). +- Determinism in this recipe comes from fixed seed + fixed channel sampling flags. +- `--determinism-mode full_greedy` is available, but can produce silent clips in some cases. +- For short prompts without references, avoid forcing a long `duration_s`; let natural stopping decide clip length. +- Keep `text`, `instruct`, `duration_s`/`tokens`, and reference inputs identical across runs for reproducibility. + +## Generation Controls + +### Common fields + +| Field | Meaning | +|---|---| +| `text` | Primary user text | +| `ref_audio`, `ref_text` | Voice/reference conditioning | +| `instruct` | Style or voice instruction | +| `tokens` | Explicit target token budget | +| `duration_s` / `seconds` | Convenience duration mapped at 12.5 tokens/sec | +| `quality` | Quality hint string (variant-dependent behavior) | +| `input_type` | `text`, `pinyin`, or `ipa` | +| `preset` | Variant sampling defaults | +| `stream`, `streaming_interval` | Chunked output controls | +| `repetition_window` | Realtime-only repetition-penalty history window | + +### `quality` hint (advisory) + +This integration treats `quality` as a user hint string. Recommended values: + +- `draft`, `balanced`, `high`, `max`, or `custom: