Dia2-2B β€” MLX, bf16

nari-labs/Dia2-2B is Nari Labs' two-speaker dialogue TTS: a script with [S1] / [S2] turns renders as one take in which both speakers trade lines. This repo is that checkpoint converted for MLX-Swift. It is the weights repo of mlx-dia2-tts-swift, a Swift port of upstream's runtime and an MLXEngine tts package.

file what dtype
model.safetensors the Dia2 decoder (28 Γ— 2048, GQA 16/8) + depformer (4 Γ— 1024, 5 weight sets) β€” upstream's keys, unchanged bf16 (RMSNorm weights float32 at load)
mimi.safetensors Kyutai's Mimi codec (32 codebooks, 12.5 Hz, 24 kHz), re-keyed onto the module layout of kyutai-labs/moshi-swift float32
config.json, tokenizer files upstream's config (+ "mimi": {"num_codebooks": 32}) and its GPT-2 BPE tokenizer β€”

Layout note. model.safetensors keeps upstream's keys. mimi.safetensors is in moshi-swift's layout: fused attention in_proj with q/k back in Kyutai's interleaved RoPE order, and conv weights as MLX (O, K, I). Loaders expecting transformers' MimiModel keys will not read it.

Conversion

Sources:

  • nari-labs/Dia2-2B @ 7abae125471a73b0fc6b9d413cb15f4ae1e771d8
  • kyutai/mimi @ 89091b3e466eb6a9d11e537bf26b144f194978f7
python Tools/oracle-capture/convert.py Dia2-2B-bf16 --size 2B --dtype bfloat16   # in mlx-dia2-tts-swift

Parity and measurements

The Swift port is gated against upstream's PyTorch runtime (nari-labs/dia2 @ 8687268) in fp32. From the package's PORTING-SPEC.md / MEASUREMENTS.md:

  • Gates: the tokenizer is id-exact; the decoder is within 1.05e-5. Every one of 4 257 sampled distributions matches upstream within total variation 1.1e-5, except one documented near-tie at the CFG filter's boundary. Codes are token-exact under replay; the waveform is within 114–118 dB.
  • This tier:
    • On upstream's own evaluation jobs it scores within sampling noise of upstream's fp32 run (scenes, single lines, voice prefixes).
    • On an M5 Max the package declares 4.5 GB resident + 3.0 GB activation (measured phys_footprint, including MLXEngine's 2 GiB buffer pool). The model's own working set is ≀ 1.2 GB, flat over takes of 3–101 s. It runs at β‰ˆ 1.8Γ— realtime.

Use (Swift)

import MLXServeCore
import MLXDia2TTS

let engine = MLXServeEngine()
try await engine.register(Dia2TTSPackage.registration, configuration: Dia2TTSConfiguration())
try await engine.prepare(.tts)
let wav = try await engine.run(TTSRequest(
    text: "[S1] Did you hear that? [S2] Hear what? It's the wind. [S1] No, it was a voice.",
    metaData: ["seed": .int(7)]))

Voice prefixes: voice: .referenceAudio(clip) + referenceTranscript set speaker 1. Speaker 2 is set with metaData speaker2Audio (a base64 .wav) + speaker2Transcript. Prefix both speakers for a consistent pair; a single prefix conditions only weakly.

Licence

  • Dia2: model.safetensors, config.json and the tokenizer files are nari-labs/Dia2-2B, Apache License 2.0 (LICENSE), Β© Nari Labs. Changes: cast to bfloat16, plus the "mimi" entry in config.json.
  • Mimi: mimi.safetensors is Mimi by Kyutai, from kyutai/mimi, licensed CC-BY-4.0. Changes: re-keyed to the moshi-swift layout, attention q/k un-permuted and re-fused, conv weights transposed, stored float32.
Downloads last month
-
Safetensors
Model size
2B params
Tensor type
BF16
Β·
MLX
Hardware compatibility
Log In to add your hardware

Quantized

Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support

Model tree for mlx-community/Dia2-2B-bf16

Quantized
(3)
this model