Skip to content

[Bug]: flashinfer 0.6.18 MLA decode tuning-config cache grows per distinct max KV length; TRT-LLM workaround no-ops #19853

Description

@vsabavat

Summary

With flashinfer-python==0.6.18 (current pin in requirements.txt), the MLA decode tuning-config cache in flashinfer grows by one entry per distinct max_seq_len value TRT-LLM passes, with no size limit. TRT-LLM's workaround for the related autotuner pathology no longer installs on 0.6.18 and fails without any warning. Host memory grows until every max KV length up to max_seq_len has been seen.

Where

  • flashinfer 0.6.18, flashinfer/mla/_core.py: _mla_decode_tuning_config is wrapped in an unbounded functools.cache. Its key includes the block-table width / profile sequence length.
  • TRT-LLM, tensorrt_llm/_torch/attention/backends/fmha/flashinfer_trtllm_gen.py: the MLA decode call to trtllm_batch_decode_with_kv_cache_mla passes params.max_past_kv_length as max_seq_len. That value changes from batch to batch, so each new max KV length adds a cache entry (~1.6 KB each) in eager MLA decode steps.
  • TRT-LLM's _install_flashinfer_mla_decode_tuning_config_cache (written for flashinfer 0.6.15) probes the builder with tensor_initializers=. On 0.6.18 that raises TypeError, the workaround catches it and returns, and nothing is logged.

Size

The growth is bounded by max_seq_len, at one entry per distinct max KV length, per rank:

  • 8K context: ~12.7 MiB
  • 128K context: ~203 MiB
  • 1M context (Kimi K3's 1,048,576): ~1.6 GiB per rank, so ~26 GiB of host memory across a 16-rank job.

Found while soaking Kimi K3 (MLA layers) on 16 x GB200 with host-memory tracing. Other host memory levelled off; this cache was the only per-length growth.

Possible fixes

  1. Pass a stable value for max_seq_len (the KV pool / block-table capacity) instead of the per-batch max_past_kv_length, so the config key does not change from batch to batch. This needs a check that the kernel only uses it as a bound.
  2. Update the workaround for 0.6.18: match the new builder signature, memoize on a bucketed profile sequence length, and log when the workaround cannot install instead of failing without a warning.
  3. Upstream in flashinfer: give _mla_decode_tuning_config an lru_cache(maxsize=...).

Environment

TensorRT-LLM main, flashinfer-python==0.6.18, PyTorch backend, Kimi K3 (MLA), 16 x GB200 (sm_100).

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    Customized kernels<NV>Specialized/modified CUDA kernels in TRTLLM for LLM ops, beyond standard TRT. Dev & perf.MemoryMemory utilization in TRTLLM: leak/OOM handling, footprint optimization, memory profiling.

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions