How to Debug KV Cache Expansion Issues Using `_expand_kv_cache()` in Pocket-TTS

To debug KV cache expansion in pocket-tts, inspect cache shapes before and after calling _expand_kv_cache(), verify NaN values remain only in new slots, and use LoggingMode to trace tensor operations that might inadvertently read uninitialized cache regions.

The pocket-tts library generates audio autoregressively using transformer layers that maintain a KV cache (key/value cache) to store past hidden states. When resuming generation from a cached model state, the private method _expand_kv_cache() in pocket_tts/models/tts_model.py must resize the cache from its current length back to the full target sequence length. Understanding how to debug this expansion prevents silent audio artifacts, context loss, and runtime device mismatches.

Understanding KV Cache Expansion in Pocket-TTS

During autoregressive generation, each transformer layer caches keys and values to avoid recomputing attention for past tokens. When you retrieve a cached model state—for example, after processing an initial audio prompt—the KV cache may only contain the portion needed for that prompt.

Before generating additional tokens, the cache must expand to accommodate the full target sequence length. This expansion is handled by _expand_kv_cache(), defined in pocket_tts/models/tts_model.py at lines 390–421:

def _expand_kv_cache(self, model_state: dict, sequence_length: int) -> None:
    """Expand KV cache back to full sequence_length for generation.

    When a model state is retrieved from cache with sliced KV caches,
    this method expands them back to the full size needed for generation.
    """
    for module_name, module_state in model_state.items():
        if "cache" in module_state:
            cache = module_state["cache"]
            # cache shape: [2, batch, cur_len, heads, dim_per_head]

            cur_len = cache.shape[2]
            if cur_len < sequence_length:
                # create an expanded cache filled with NaNs

                expanded = torch.full(
                    (cache.shape[0], cache.shape[1], sequence_length,
                     cache.shape[3], cache.shape[4]),
                    float("NaN"),
                    device=cache.device,
                    dtype=cache.dtype,
                )
                # copy existing entries

                expanded[:, :, :cur_len, :, :] = cache
                module_state["cache"] = expanded

The method is invoked inside _generate() (lines 718–722) to prepare the cache before the next generation step:

required_len = current_end + token_count + max_gen_len
self._expand_kv_cache(model_state, sequence_length=required_len)

Common KV Cache Expansion Bugs

Several specific failure modes occur when _expand_kv_cache() behaves unexpectedly or when subsequent code interacts improperly with the expanded cache:

  • NaN propagation – The expansion fills unused slots with float("NaN"). If an off-by-one error or incorrect indexing causes the model to read from these slots, the computation produces NaN values that propagate through the audio generation, resulting in silent output or artifacts.
  • Mismatched sequence lengths – If the sequence_length argument is smaller than the actual tokens to be generated, the cache truncates early, causing the model to lose context and produce abrupt speech cutoffs.
  • Device and dtype mismatches – The expanded tensor must reside on the same device (CPU/GPU) and use the same dtype as the original cache. While the implementation explicitly passes device=cache.device and dtype=cache.dtype, manual model manipulation elsewhere can violate this assumption.

Step-by-Step Debugging Strategy

Inspect Cache Shapes

Before and after expansion, verify that the cache dimensions match expectations. The cache tensor has shape [2, batch, seq_len, heads, dim_per_head], where the first dimension represents keys and values.

def print_cache_shapes(state):
    for name, s in state.items():
        if "cache" in s:
            print(f"{name}: {s['cache'].shape}")

print("Before expansion:")
print_cache_shapes(model_state)
self._expand_kv_cache(model_state, required_len)
print("After expansion:")
print_cache_shapes(model_state)

Validate NaN Placement

After expansion, only the newly allocated positions should contain NaN values. Verify that the original cached data remains intact:

cache = model_state["transformer"]["cache"]  # adjust key as needed

cur_len = cache.shape[2]
assert not torch.isnan(cache[:, :, :cur_len, :, :]).any(), "Existing cache corrupted"
assert torch.isnan(cache[:, :, cur_len:, :, :]).any(), "New slots not filled with NaN"

Trace Operations with LoggingMode

To detect which operation reads from the NaN region, use the LoggingMode utility from pocket_tts/utils/debugging.py. This context manager prints each Aten call with arguments and outputs, revealing accidental cache access:

from pocket_tts.utils.debugging import LoggingMode

with LoggingMode():
    self._run_flow_lm_and_increment_step(model_state, ...)

Verify Device and Dtype Consistency

If you manually moved the model to GPU, confirm the expansion respects the cache's device placement:

for name, s in model_state.items():
    if "cache" in s:
        cache = s["cache"]
        print(f"{name}: device={cache.device}, dtype={cache.dtype}")

Check Sequence Length Calculations

The required_len calculation in _generate() combines several variables. Print these components to ensure the expansion size is correct:

print(f"current_end: {current_end}")
print(f"token_count: {token_count}")
print(f"max_gen_len: {max_gen_len}")
print(f"required_len: {required_len}")

Complete Debugging Example

The following snippet demonstrates a full debugging workflow that loads a model, inspects cache states, runs generation with logging enabled, and verifies post-generation cache integrity:

from pocket_tts import TTSModel
from pocket_tts.utils.debugging import LoggingMode
import torch

model = TTSModel.load_model()
voice_state = model.get_state_for_audio_prompt(
    "hf://kyutai/tts-voices/alba-mackenna/casual.wav"
)

text = "Hello world!"

def dump_cache(state, label):
    print(f"\n[{label}] KV-cache shapes:")
    for name, s in state.items():
        if "cache" in s:
            cache = s["cache"]
            print(f"  {name}: shape={cache.shape}, device={cache.device}")

def verify_cache_integrity(state, original_len):
    """Check that only new positions contain NaN."""
    for name, s in state.items():
        if "cache" in s:
            cache = s["cache"]
            existing = cache[:, :, :original_len, :, :]
            new = cache[:, :, original_len:, :, :]
            assert not torch.isnan(existing).any(), f"{name}: Existing values erased"
            assert torch.isnan(new).any(), f"{name}: New positions not NaN"

# Capture initial state

dump_cache(voice_state, "initial")
initial_len = voice_state["transformer"]["cache"].shape[2]

# Run generation with tracing

with LoggingMode():
    audio = model.generate_audio(
        voice_state,
        text,
        frames_after_eos=2,
        copy_state=False,  # modify in-place to observe expansion

    )

# Verify expansion

dump_cache(voice_state, "post-generation")
verify_cache_integrity(voice_state, initial_len)

Key Source Files to Reference

File Role
pocket_tts/models/tts_model.py Implements _expand_kv_cache() and the generation pipeline.
pocket_tts/modules/stateful_module.py Base class for streaming modules that store the KV cache.
pocket_tts/utils/debugging.py Provides LoggingMode for low-level Torch dispatch debugging.

Summary

  • _expand_kv_cache() in pocket_tts/models/tts_model.py resizes KV caches by filling new positions with NaN while preserving existing values.
  • NaN propagation is the primary symptom of bugs, caused by reading uninitialized slots after expansion.
  • Debug by comparing shapes before and after expansion using the print helper functions.
  • Use LoggingMode from pocket_tts/utils/debugging.py to trace which operations touch the cache memory.
  • Validate device and dtype consistency to prevent runtime errors when using GPU acceleration.
  • Check length calculations (current_end + token_count + max_gen_len) when audio cuts off abruptly.

Frequently Asked Questions

What causes NaN values in the KV cache during generation?

NaN values appear intentionally in the expanded region of the KV cache to mark uninitialized positions. However, if your generated audio contains NaNs or silence, it typically indicates that the model is reading from these uninitialized slots due to an off-by-one indexing error or incorrect attention masking after _expand_kv_cache() resizes the cache.

How does _expand_kv_cache() handle device placement?

The method explicitly creates the expanded tensor using device=cache.device and dtype=cache.dtype to ensure the new memory resides on the same device as the original cache. If you manually move tensors between CPU and GPU outside the model's standard flow, verify that model_state remains on the expected device before expansion.

Why is my generated audio silent after cache expansion?

Silent output usually results from NaN propagation through the attention mechanism. After calling _expand_kv_cache(), confirm that only the newly allocated temporal positions (indices greater than the original cur_len) contain NaN values. If existing positions were overwritten or if subsequent layers read beyond the valid cache length, the attention scores become NaN, silencing the audio output.

How can I verify the cache expansion worked correctly?

Validate the expansion by checking three properties: (1) the cache shape increased from [2, batch, old_len, heads, dim] to [2, batch, new_len, heads, dim]; (2) torch.isnan(cache[:, :, old_len:, :, :]).any() returns True for the new region only; and (3) torch.isnan(cache[:, :, :old_len, :, :]).any() returns False, confirming original data preservation.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →