Instructions to use OpenMOSS-Team/MOSS-Audio-Tokenizer-Nano with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use OpenMOSS-Team/MOSS-Audio-Tokenizer-Nano with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="OpenMOSS-Team/MOSS-Audio-Tokenizer-Nano", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("OpenMOSS-Team/MOSS-Audio-Tokenizer-Nano", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Fix streaming decode with use_cuda_graph=True: first frame never runs, K-frame calls repeat the last frame
Hi! batch_decode(streaming=True, use_cuda_graph=True) gives wrong audio for me. I traced it to _graphed_decode_frame:
- On the first call the graph is captured but not replayed. Capture doesn't run anything, so frame 0 is never decoded and the decoder state skips it.
- It returns the graph's static output buffers.
step()keeps a view per frame and only concatenates after the loop, so when a call decodes more than one frame, all of them end up as the last one.
This PR replays right after capture and returns copies. The non-graph path is unchanged.
Tested on an RTX 4070 Ti SUPER (torch 2.8, fp32), 64 frames of random codes, 4 frames per call, compared with use_cuda_graph=False:
| frames out | frames matching plain decode | max abs diff | ms / frame | |
|---|---|---|---|---|
| plain decode | 64 | – | – | 19.2 |
| graph, before | 63 | 0% | – | 3.8 |
| graph, after | 64 | 100% | 0.0 | 3.7 |
The fixed version also matches exactly with 1 and 3 frames per call, and the same patch checked out on a Jetson AGX Orin.
For context, we hit this through VieNeu-TTS v3 Turbo, which keeps the graph off because of it. With the fix and the graph on, its end-to-end RTF went from 0.84 to 0.34 on the Orin and from 0.29 to 0.15 on the 4070 Ti SUPER, with no change in WER.
I only tested the SDPA path (no flash_attn installed).
Repro script
"""Repro: streaming batch_decode(use_cuda_graph=True) returns wrong audio.
Decodes the same random codes (as in the README example) with
`batch_decode(streaming=True)`, K frames per call, with `use_cuda_graph=False`
as the reference and `use_cuda_graph=True` before and after the proposed fix,
then compares the output frame by frame (one frame = `downsample_rate` samples).
python repro_cuda_graph.py [--frames 64] [--seed 0]
The patched rows wrap MossAudioTokenizerDecodeSession._graphed_decode_frame at
runtime with the two changes of the PR, one at a time and together: replay right
after capture, and return copies of the static output buffers.
"""
import argparse
import sys
import time
import numpy as np
import torch
from transformers import AutoModel
REPO = "OpenMOSS-Team/MOSS-Audio-Tokenizer-Nano"
def decode(model, codes, k, graph):
"""Stream-decode `codes` (n_q, T) k frames per call. Returns (audio (C, N), ms per frame)."""
model._reset_batch_decode_streaming_state()
chunks, per_frame, total = [], [], codes.shape[1]
for i, lo in enumerate(range(0, total, k)):
block = codes[:, lo:lo + k]
torch.cuda.synchronize()
t0 = time.perf_counter()
out = model.batch_decode(
[block], streaming=True, max_batch_size=1, reset_stream=(i == 0),
finalize_indices=[0] if lo + k >= total else None, use_cuda_graph=graph,
)
torch.cuda.synchronize()
if i > 0: # the first call of a graph session includes the capture
per_frame.append((time.perf_counter() - t0) / block.shape[1])
chunks.append(out.audio[0, :, :int(out.audio_lengths[0])].float().cpu())
model._reset_batch_decode_streaming_state()
return torch.cat(chunks, dim=-1).numpy(), 1000 * float(np.mean(per_frame))
def compare(ref, out, hop):
"""Frames of `out` matching `ref` (corr > 0.99) at lag 0 and at lag 1 (out[k] vs ref[k+1])."""
def frames(x):
return [x[:, j * hop:(j + 1) * hop].ravel() for j in range(x.shape[-1] // hop)]
fr, fo = frames(ref), frames(out)
res = {"frames": f"{len(fo)} / {len(fr)}"}
for lag in (0, 1):
pairs = [(fr[j + lag], fo[j]) for j in range(min(len(fo), len(fr) - lag))]
ok = [j + lag for j, (a, b) in enumerate(pairs) if np.corrcoef(a, b)[0, 1] > 0.99]
res[f"lag{lag}"] = f"{len(ok) / max(len(pairs), 1):.0%} {ok[:4]}{'…' if len(ok) > 4 else ''}"
res["max |diff|"] = f"{np.abs(ref - out).max():.1e}" if ref.shape == out.shape else "-"
return res
FIX = {"replay": False, "clone": False}
def install_fix(model):
"""Wrap _graphed_decode_frame; FIX selects which of the two changes are active."""
session = sys.modules[type(model).__module__].MossAudioTokenizerDecodeSession
graphed = session._graphed_decode_frame
def patched(self, codes, code_lengths):
key = (str(codes.device), self.max_batch_size, codes.shape[0], self.model.compute_dtype_name)
capturing = self._cuda_graph is None or self._cuda_graph_key != key
out = graphed(self, codes, code_lengths)
if capturing and FIX["replay"]:
self._cuda_graph.replay() # capture only records the kernels
if FIX["clone"]:
return type(out)(audio=out.audio.clone(), audio_lengths=out.audio_lengths.clone())
return out
session._graphed_decode_frame = patched # with both flags off: original behaviour
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--frames", type=int, default=64)
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
model = AutoModel.from_pretrained(REPO, trust_remote_code=True).eval().to("cuda")
q = model.config.quantizer_kwargs
hop = model.config.downsample_rate
torch.manual_seed(args.seed)
codes = torch.randint(0, q["codebook_size"], (q["num_quantizers"], args.frames), device="cuda")
print(f"{torch.cuda.get_device_name()} · torch {torch.__version__} · "
f"{args.frames} frames · compute {model.compute_dtype_name}\n")
install_fix(model)
cases = [("no graph (reference)", False, None),
("graph, original", True, (False, False)),
("graph, only replay-after-capture", True, (True, False)),
("graph, only clone outputs", True, (False, True)),
("graph, both (this PR)", True, (True, True))]
rows = []
with torch.no_grad():
for k in (4, 1):
ref = None
for name, graph, fix in cases:
if fix is not None:
FIX.update(replay=fix[0], clone=fix[1])
out, ms = decode(model, codes, k, graph)
ref = out if ref is None else ref
rows.append((f"K={k} · {name}", ms, compare(ref, out, hop)))
print("| decode | ms / frame | frames out / ref | match at lag 0 | match at lag 1 | max abs diff |")
print("|---|---:|---|---|---|---|")
for name, ms, r in rows:
print(f"| {name} | {ms:.1f} | {r['frames']} | {r['lag0']} | {r['lag1']} | {r['max |diff|']} |")
if __name__ == "__main__":
main()