Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions fish_speech/models/text2semantic/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,12 +104,14 @@ def decode_one_token_ar(
audio_masks: torch.Tensor,
audio_parts: torch.Tensor,
previous_tokens: Optional[torch.Tensor] = None,
kv_len: Optional[int] = None,
) -> torch.Tensor:
forward_result = model.forward_generate(
x,
input_pos,
audio_masks=audio_masks,
audio_parts=audio_parts,
kv_len=kv_len,
)
logits = forward_result.logits # (1, 1, vocab_size)
hidden_states = forward_result.hidden_states
Expand Down Expand Up @@ -193,7 +195,13 @@ def decode_n_tokens(
audio_masks: torch.Tensor,
audio_parts: torch.Tensor,
decode_one_token=decode_one_token_ar,
kv_start_pos: Optional[int] = None,
):
if kv_start_pos is None:
# Compatibility fallback for direct callers. The production generation
# path passes the Python position explicitly to avoid a device sync.
kv_start_pos = int(input_pos[0].item())

# Rolling window for RAS (Repetition Aware Sampling)
previous_tokens = torch.zeros(
(model.config.num_codebooks + 1, RAS_WIN_SIZE),
Expand All @@ -212,6 +220,7 @@ def decode_n_tokens(
model=model,
x=cur_token,
input_pos=input_pos,
kv_len=kv_start_pos + i + 1,
previous_tokens=previous_tokens,
temperature=temperature,
top_p=top_p,
Expand Down Expand Up @@ -331,6 +340,7 @@ def generate(
semantic_logit_bias,
audio_masks,
audio_parts,
kv_len=T,
)
seq[:, T : T + 1] = first_token

Expand All @@ -349,6 +359,7 @@ def generate(
audio_masks=audio_masks,
audio_parts=audio_parts,
decode_one_token=decode_one_token,
kv_start_pos=T,
)
seq = seq[:, : T + 1 + x.size(1)]
seq[:, T + 1 :] = x
Expand Down
39 changes: 33 additions & 6 deletions fish_speech/models/text2semantic/llama.py
Original file line number Diff line number Diff line change
Expand Up @@ -394,6 +394,7 @@ def forward_generate(
audio_masks: Optional[Tensor] = None,
audio_parts: Optional[Tensor] = None,
return_all: bool = False,
kv_len: Optional[int] = None,
) -> BaseTransformerForwardResult:

# Embedding logic replicated from embed() for compilation compatibility
Expand Down Expand Up @@ -432,13 +433,23 @@ def forward_generate(
else:
logger.warning("audio_parts provided but model has no audio_projector")

cache_capacity = (
self.max_seq_len if self.max_seq_len > 0 else self.config.max_seq_len
)
if input_pos is None:
input_pos = torch.arange(inp.shape[-1], device=x.device)
max_seq_len = inp.shape[-1]
active_kv_len = inp.shape[-1]
else:
max_seq_len = self.max_seq_len
active_kv_len = cache_capacity if kv_len is None else kv_len

if not 1 <= active_kv_len <= cache_capacity:
raise ValueError(
f"kv_len must be between 1 and {cache_capacity}, got {active_kv_len}"
)

mask = self.causal_mask[None, None, input_pos, :max_seq_len] # (B, N, Q, K)
mask = self.causal_mask[
None, None, input_pos, :active_kv_len
] # (B, N, Q, active K)
freqs_cis = self.freqs_cis[input_pos]

for layer in self.layers:
Expand Down Expand Up @@ -651,9 +662,12 @@ def forward(
return self.decode(result)

def forward_generate(
self, x: Tensor, input_pos: Optional[Tensor] = None
self,
x: Tensor,
input_pos: Optional[Tensor] = None,
kv_len: Optional[int] = None,
) -> TransformerForwardResult:
result = super().forward_generate(x, input_pos)
result = super().forward_generate(x, input_pos, kv_len=kv_len)
return self.decode(result)


Expand Down Expand Up @@ -822,8 +836,15 @@ def forward_generate(
input_pos: Optional[Tensor] = None,
audio_masks: Optional[Tensor] = None,
audio_parts: Optional[Tensor] = None,
kv_len: Optional[int] = None,
) -> TransformerForwardResult:
x = super().forward_generate(x, input_pos, audio_masks, audio_parts)
x = super().forward_generate(
x,
input_pos,
audio_masks,
audio_parts,
kv_len=kv_len,
)
x.hidden_states = self.fast_project_in(x.hidden_states)
return x

Expand Down Expand Up @@ -909,6 +930,12 @@ def forward(

if self.kv_cache is not None:
k, v = self.kv_cache.update(input_pos, k, v)
if mask is not None:
# Keep the cache allocated at max_seq_len, but repeat and
# attend only to the prefix populated so far.
active_kv_len = mask.shape[-1]
k = k[:, :, :active_kv_len]
v = v[:, :, :active_kv_len]

k = k.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
v = v.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
Expand Down