Skip to content
Merged
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
154 changes: 110 additions & 44 deletions tensorrt_llm/_torch/attention_backend/fmha/flashinfer_trtllm_gen.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,35 +67,49 @@
)


def _patch_flashinfer_dynamic_tensor_spec_eq() -> None:
"""Work around a hash/eq inconsistency in flashinfer 0.6.15's autotuner.

``DynamicTensorSpec`` defines a custom ``__hash__`` that skips
``tensor_initializers`` (and hashes callables by identity), but keeps the
dataclass-generated ``__eq__``, which compares all fields by value.
Callers such as ``trtllm_batch_decode_with_kv_cache_mla`` build a fresh
``TuningConfig`` per call with fresh initializer closures, so the
def _install_flashinfer_mla_decode_tuning_config_cache() -> None:
"""Work around an autotuner cache pathology on flashinfer 0.6.15's MLA decode path.

``flashinfer.mla._core._build_mla_decode_tuning_config`` builds a fresh
``TuningConfig`` — with fresh initializer closures — on every
``trtllm_batch_decode_with_kv_cache_mla`` call. ``DynamicTensorSpec``'s
custom ``__hash__`` skips ``tensor_initializers``, but its
dataclass-generated ``__eq__`` compares them by value, so each call's
config is hash-equal but eq-unequal to all previous ones: the
``lru_cache`` on ``AutoTuner._find_nearest_profile`` accumulates
hash-equal but eq-unequal keys up to the cache cap; every autotuner cache
probe on the MLA decode path then walks the entire collision chain in
Python ``__eq__`` (tens of milliseconds per attention call, per layer).

The workaround replaces ``__eq__`` with one consistent with the existing
``__hash__``: ignore ``tensor_initializers``, compare
``map_to_tuning_buckets`` by identity, and ``gen_tuning_buckets`` by
value if it is a tuple and by identity otherwise. Keys equal under this
``__eq__`` produce identical nearest profiles, so cached results stay
correct, and hash/eq consistency is restored (equal objects hash equal).
colliding keys up to the cache cap, and every autotuner cache probe on
the MLA decode path then walks the entire collision chain in Python
``__eq__`` (tens of milliseconds per attention call, per layer).

The workaround memoizes the config builder on the scalars that fully
determine its output (kv-cache paging geometry, workspace size, runner
set, head/rank dims, max_seq_len, device), so repeated calls reuse the
same ``TuningConfig`` object. Cache hits then resolve through
flashinfer's own unmodified hash/eq — no equality semantics change
anywhere, and the scope is exactly the MLA decode path that exhibits the
problem.

MLA decode is the only path that needs this: every other ``TuningConfig``
construction site in flashinfer 0.6.15 already reuses its config across
calls (module-level constants in ``gemm``, instance-level configs in the
cute-dsl MoE tuner, ``functools.cache``-wrapped builders in
``mla/_sparse_mla_sm120.py``), so those specs hit the autotuner caches by
object identity and never accumulate collision chains. Memoizing the one
per-call builder therefore covers the whole pathology; extending it to
other ops would add nothing.

The bug is upstream in flashinfer; the patch is applied only when a
behavioral probe shows the inconsistency (two specs differing only in
initializer closures that hash equal but compare unequal), so a fixed
flashinfer release degrades this to a no-op.
behavioral probe shows the hash/eq inconsistency and the builder still
has the expected signature, so a fixed or refactored flashinfer release
degrades this to a no-op.
"""
try:
import inspect

from flashinfer.autotuner.autotuner import DynamicTensorSpec
from flashinfer.mla import _core as _flashinfer_mla_core

def _probe_spec():
def _probe_spec() -> DynamicTensorSpec:
return DynamicTensorSpec(
input_idx=(0,),
dim_idx=(0,),
Expand All @@ -108,37 +122,89 @@ def _probe_spec():
if not (hash(a) == hash(b) and a != b):
return # flashinfer without the inconsistency; nothing to do

def _hash_consistent_eq(self, other):
if self is other:
return True
if not isinstance(other, DynamicTensorSpec):
return NotImplemented
if isinstance(self.gen_tuning_buckets, tuple) and isinstance(
other.gen_tuning_buckets, tuple
):
buckets_equal = self.gen_tuning_buckets == other.gen_tuning_buckets
else:
buckets_equal = self.gen_tuning_buckets is other.gen_tuning_buckets
return (
self.input_idx == other.input_idx
and self.dim_idx == other.dim_idx
and buckets_equal
and self.map_to_tuning_buckets is other.map_to_tuning_buckets
orig_build = _flashinfer_mla_core._build_mla_decode_tuning_config
if getattr(orig_build, "_trtllm_mla_tuning_config_cache", False):
return # already installed

expected_params = (
"kv_cache",
"block_tables",
"workspace_buffer",
"runner_names",
"q_len",
"num_heads",
"kv_lora_rank",
"max_seq_len",
"device",
)
signature = inspect.signature(orig_build)
if tuple(signature.parameters) != expected_params:
return # builder was refactored; memo key may no longer be complete
try:
signature.bind(**dict.fromkeys(expected_params))
except TypeError:
return # parameters went positional-only; the keyword call below would fail

cache: dict = {}

def _cached_build(
kv_cache,
block_tables,
workspace_buffer,
runner_names,
q_len,
num_heads,
kv_lora_rank,
max_seq_len,
device,
):
# Everything the builder (and _compute_mla_decode_buckets) reads
# from its tensor arguments reduces to these scalars.
key = (
kv_cache.shape[0], # num_pages, captured by init_block_tables
kv_cache.shape[-2], # page_size, feeds profile_seq_len
block_tables.shape[-1], # pages per sequence, feeds profile_seq_len
workspace_buffer.numel() * workspace_buffer.element_size(), # cute-dsl bucket cap
tuple(runner_names),
q_len,
num_heads,
kv_lora_rank,
max_seq_len,
device,
)
config = cache.get(key)
if config is None:
if len(cache) >= 256: # backstop; a serving process sees a handful of keys
cache.clear()
config = orig_build(
kv_cache=kv_cache,
block_tables=block_tables,
workspace_buffer=workspace_buffer,
runner_names=runner_names,
q_len=q_len,
num_heads=num_heads,
kv_lora_rank=kv_lora_rank,
max_seq_len=max_seq_len,
device=device,
)
cache[key] = config
return config

DynamicTensorSpec.__eq__ = _hash_consistent_eq
_cached_build._trtllm_mla_tuning_config_cache = True
_flashinfer_mla_core._build_mla_decode_tuning_config = _cached_build
logger.debug(
"Patched flashinfer DynamicTensorSpec.__eq__ to be hash-consistent "
"Memoized flashinfer _build_mla_decode_tuning_config "
"(autotuner cache collision-chain workaround)."
)
except Exception:
# A future flashinfer refactor (moved class, changed fields) must not
except (ImportError, AttributeError, TypeError, ValueError):
# A future flashinfer refactor (moved module, changed fields) must not
# break attention; it just loses the workaround.
logger.debug("Skipping flashinfer DynamicTensorSpec.__eq__ workaround.", exc_info=True)
# tensorrt_llm's logger.debug(*msg) takes no exc_info/kwargs.
logger.debug("Skipping flashinfer MLA decode tuning-config cache workaround.")


if IS_FLASHINFER_AVAILABLE:
_patch_flashinfer_dynamic_tensor_spec_eq()
_install_flashinfer_mla_decode_tuning_config_cache()


_SUPPORTED_MLA_BACKENDS = {"cute-dsl", "trtllm-gen"}
Expand Down
Loading