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

#3
Files changed (1) hide show
  1. modeling_moss_audio_tokenizer.py +7 -4
modeling_moss_audio_tokenizer.py CHANGED
@@ -560,12 +560,15 @@ class MossAudioTokenizerDecodeSession:
560
  self._cuda_graph_key = cuda_graph_key
561
  self._graph_output_audio = decoder_output.audio
562
  self._graph_output_audio_lengths = decoder_output.audio_lengths
563
- else:
564
- self._cuda_graph.replay()
565
 
 
 
 
 
 
566
  return MossAudioTokenizerDecoderOutput(
567
- audio=self._graph_output_audio,
568
- audio_lengths=self._graph_output_audio_lengths,
569
  )
570
 
571
  def _reset_slot(self, slot_index: int) -> None:
 
560
  self._cuda_graph_key = cuda_graph_key
561
  self._graph_output_audio = decoder_output.audio
562
  self._graph_output_audio_lengths = decoder_output.audio_lengths
 
 
563
 
564
+ # Capture only records the kernels, so run the graph for this frame too.
565
+ self._cuda_graph.replay()
566
+
567
+ # Every replay overwrites these static buffers, while step() keeps each
568
+ # frame's output until the end of the call: hand out copies.
569
  return MossAudioTokenizerDecoderOutput(
570
+ audio=self._graph_output_audio.clone(),
571
+ audio_lengths=self._graph_output_audio_lengths.clone(),
572
  )
573
 
574
  def _reset_slot(self, slot_index: int) -> None: