Fix streaming decode with use_cuda_graph=True: first frame never runs, K-frame calls repeat the last frame

#3
by noland0520 - opened

Hi! batch_decode(streaming=True, use_cuda_graph=True) gives wrong audio for me. I traced it to _graphed_decode_frame:

  1. 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.
  2. 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()
Ready to merge
This branch is ready to get merged automatically.

Sign up or log in to comment