Fix rotary/kvcache OOB with speaker_audio + CUDA/CPU device mix

#13
by multimodalart HF Staff - opened
Files changed (1) hide show
  1. zonos/model.py +3 -2
zonos/model.py CHANGED
@@ -199,7 +199,8 @@ class Zonos(nn.Module):
199
  audio_seq_len = prefix_audio_len + max_new_tokens
200
  seq_len = prefix_conditioning.shape[1] + audio_seq_len
201
 
202
- inference_params = self.setup_cache(batch_size=batch_size * 2, max_seqlen=seq_len)
 
203
 
204
  codes = torch.full((batch_size, 9, audio_seq_len), unknown_token, device="cuda")
205
  if audio_prefix_codes is not None:
@@ -237,7 +238,7 @@ class Zonos(nn.Module):
237
  next_token = sample_from_logits(logits, generated_tokens=delayed_codes[..., :offset], **sampling_params)
238
  eos_in_cb0 = next_token[:, 0] == self.eos_token_id
239
 
240
- remaining_steps[eos_in_cb0[:, 0]] = torch.minimum(remaining_steps[eos_in_cb0[:, 0]], torch.tensor(9))
241
  stopping |= eos_in_cb0[:, 0]
242
 
243
  eos_codebook_idx = 9 - remaining_steps
 
199
  audio_seq_len = prefix_audio_len + max_new_tokens
200
  seq_len = prefix_conditioning.shape[1] + audio_seq_len
201
 
202
+ # +9 headroom for the delay-pattern tail (delayed_codes is 9 longer than audio_seq_len)
203
+ inference_params = self.setup_cache(batch_size=batch_size * 2, max_seqlen=seq_len + 9)
204
 
205
  codes = torch.full((batch_size, 9, audio_seq_len), unknown_token, device="cuda")
206
  if audio_prefix_codes is not None:
 
238
  next_token = sample_from_logits(logits, generated_tokens=delayed_codes[..., :offset], **sampling_params)
239
  eos_in_cb0 = next_token[:, 0] == self.eos_token_id
240
 
241
+ remaining_steps[eos_in_cb0[:, 0]] = remaining_steps[eos_in_cb0[:, 0]].clamp(max=9)
242
  stopping |= eos_in_cb0[:, 0]
243
 
244
  eos_codebook_idx = 9 - remaining_steps