Skip to content
Merged
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
29 changes: 20 additions & 9 deletions tensorrt_llm/_torch/attention_backend/flashinfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -663,17 +663,28 @@ def _post_init_with_buffers(self, buffers) -> None:
max_num_blocks = max(max_num_blocks, lbuf.shape[0])
self._max_num_blocks = max_num_blocks

# Detect VSWA: check if the manager has multiple pools.
# Guard on layer_to_pool_mapping_dict which is V2-specific — V1
# managers also expose is_vswa but lack the per-pool infrastructure.
if (getattr(self.kv_cache_manager, 'is_vswa', False) and hasattr(
self.kv_cache_manager, 'layer_to_pool_mapping_dict')):
mgr = self.kv_cache_manager
self._vswa_layer_to_pool = {}
self._vswa_pool_to_rep_layer: Dict[int, int] = {}
# Layers may share one page-index list only when they are in the
# same pool AND have the same page-index scale: VSWA splits pools,
# and per-layer geometry (e.g. Gemma4 sliding/global head_dim)
# splits scales even when all windows collapse to max_seq_len and
# is_vswa is False. Guarded on V2-specific attributes (V1 managers
# lack the per-pool infrastructure).
mgr = self.kv_cache_manager
get_scale = getattr(mgr, 'get_layer_page_index_scale', None)
layer_space: Dict[int, int] = {}
if hasattr(mgr, 'layer_to_pool_mapping_dict') and get_scale:
space_ids = {}
for layer_idx in getattr(mgr, 'layer_offsets', {}):
layer_offset = mgr.layer_offsets[layer_idx]
pool_id = mgr.layer_to_pool_mapping_dict[layer_offset]
key = (mgr.layer_to_pool_mapping_dict[layer_offset],
get_scale(layer_idx))
space_ids.setdefault(key, len(space_ids))
layer_space[layer_idx] = space_ids[key]
if layer_space and (getattr(mgr, 'is_vswa', False)
or len(set(layer_space.values())) > 1):
self._vswa_layer_to_pool = {}
self._vswa_pool_to_rep_layer: Dict[int, int] = {}
for layer_idx, pool_id in layer_space.items():
self._vswa_layer_to_pool[layer_idx] = pool_id
if pool_id not in self._vswa_pool_to_rep_layer:
self._vswa_pool_to_rep_layer[pool_id] = layer_idx
Expand Down
21 changes: 18 additions & 3 deletions tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -2796,6 +2796,12 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False):
else:
self.index_mapper.remove_sequence(request.py_request_id)

def get_layer_page_index_scale(self, layer_idx: int) -> int:
"""Page-index scale of this layer's KV buffer. Layers in one pool can
have different scales (e.g. different head_dim), so per-layer callers
must not use the pool-level scale."""
return int(self.impl.get_page_index_scale(self.layer_offsets[layer_idx], Role.KEY))

def get_batch_cache_indices(
self,
request_ids: List[int],
Expand All @@ -2804,13 +2810,16 @@ def get_batch_cache_indices(
) -> List[List[int]]:
if layer_idx is None:
pool_id = 0
index_scale = None
else:
pool_id = self.layer_to_pool_mapping_dict[self.layer_offsets[layer_idx]]
index_scale = self.get_layer_page_index_scale(layer_idx)
return self._get_batch_cache_indices_by_pool_id(
request_ids,
pool_id=pool_id,
is_kv_aggregate=True,
num_blocks_per_seq=num_blocks_per_seq,
index_scale=index_scale,
)

def _get_batch_cache_indices_by_pool_id(
Expand All @@ -2820,6 +2829,7 @@ def _get_batch_cache_indices_by_pool_id(
pool_id: int = 0,
is_kv_aggregate: bool = True,
num_blocks_per_seq: Optional[Sequence[int]] = None,
index_scale: Optional[int] = None,
) -> List[List[int]]:
if is_kv_aggregate:
# Div by kv_factor to index kv cache with size
Expand All @@ -2828,7 +2838,8 @@ def _get_batch_cache_indices_by_pool_id(
else:
div_factor = 1

index_scale = int(self.index_scales[pool_id])
if index_scale is None:
index_scale = int(self.index_scales[pool_id])
res = []

for req_idx, req_id in enumerate(request_ids):
Expand Down Expand Up @@ -2871,10 +2882,14 @@ def get_batch_cache_indices_flat(
"""
if layer_idx is None:
pool_id = 0
scale = self._index_scale_ints[pool_id]
else:
pool_id = self.layer_to_pool_mapping_dict[self.layer_offsets[layer_idx]]

scale = self._index_scale_ints[pool_id]
# Layers sharing a pool can still require different page-index
# scales (e.g. Gemma4 sliding/global head_dim). Use the per-layer
# scale so this flat block table matches get_batch_cache_indices()
# and never feeds out-of-range page ids to FlashInfer.
scale = self.get_layer_page_index_scale(layer_idx)
div_factor = self.kv_factor

out_tensor = torch.empty(sum(num_blocks), dtype=torch.int32, pin_memory=prefer_pinned())
Expand Down
2 changes: 2 additions & 0 deletions tests/integration/test_lists/test-db/l0_b200.yml
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,8 @@ l0_b200:
- unittest/_torch/modeling/test_gemma4_multimodal.py
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_text_26b_dummy
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_text_e2b_dummy
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_text_e2b_dummy_small_max_seq_len[256]
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_text_e2b_dummy_small_max_seq_len[512]
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_text_31b_dummy
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_text_e4b_dummy
- unittest/_torch/modeling/test_gemma4_e2e_dummy.py::test_e2e_multimodal_26b_dummy
Expand Down
126 changes: 99 additions & 27 deletions tests/unittest/_torch/modeling/test_gemma4_e2e_dummy.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,19 @@
_LLM_MODELS_ROOT = os.environ.get("LLM_MODELS_ROOT")
if _LLM_MODELS_ROOT is None:
pytest.skip("LLM_MODELS_ROOT not set", allow_module_level=True)
_GEMMA4_MODELS = os.path.join(_LLM_MODELS_ROOT, "gemma4")
# Canonical model root subdir is "gemma" (see tests/test_common/llm_data.py:
# "google/gemma-4-E2B-it" -> "gemma/gemma-4-E2B-it"), not "gemma4".
_GEMMA4_MODELS = os.path.join(_LLM_MODELS_ROOT, "gemma")

# Imported after the module-level skip guard so that collecting this module on
# a machine without LLM_MODELS_ROOT does not pull in the runtime import.
from tensorrt_llm.llmapi import LLM, KvCacheConfig, SamplingParams # noqa: E402

# These dummy models are tiny, but the default KV-cache fraction sizes the pool
# to most of the (very large B200) device memory, leaving nothing for the other
# executor components -> OOM at executor creation. Cap it so the whole pipeline
# fits regardless of card size.
_KV_CACHE_CONFIG = KvCacheConfig(free_gpu_memory_fraction=0.5)

# Real model paths — used for tokenizer + base config
MODEL_PATHS = {
Expand All @@ -53,11 +65,6 @@
}


def _model_available(name: str) -> bool:
path = MODEL_PATHS.get(name, "")
return os.path.isfile(os.path.join(path, "config.json"))


def _make_dummy_config_dir(
model_path: str,
dummy_head_dim: int = 128,
Expand All @@ -73,10 +80,20 @@ def _make_dummy_config_dir(

Returns path to the temp directory.
"""
# Fail loudly rather than silently skipping when the real checkpoint is
# missing: these tests are registered in CI with LLM_MODELS_ROOT set, so an
# absent artifact is a setup error, not a reason to report a passing skip.
config_path = os.path.join(model_path, "config.json")
if not os.path.isfile(config_path):
raise FileNotFoundError(
f"Gemma4 test checkpoint not found: {config_path}. These E2E tests "
f"require the real gemma-4 models under $LLM_MODELS_ROOT/gemma."
)

tmp_dir = tempfile.mkdtemp(prefix="gemma4_dummy_")

# Load and patch config
with open(os.path.join(model_path, "config.json")) as f:
with open(config_path) as f:
config = json.load(f)

tc = config.get("text_config", config)
Expand Down Expand Up @@ -176,14 +193,17 @@ def _make_dummy_config_dir(


@requires_gemma4_transformers
@pytest.mark.skipif(not _model_available("26B"), reason="gemma-4-26B-A4B-it not found")
def test_e2e_text_26b_dummy():
"""E2E text generation for 26B-A4B (MoE + K=V + softcap + hybrid attn)."""
from tensorrt_llm.llmapi import LLM, SamplingParams

dummy_dir = _make_dummy_config_dir(MODEL_PATHS["26B"])
try:
llm = LLM(dummy_dir, load_format="dummy", attn_backend="FLASHINFER", dtype="bfloat16")
llm = LLM(
dummy_dir,
load_format="dummy",
attn_backend="FLASHINFER",
dtype="bfloat16",
kv_cache_config=_KV_CACHE_CONFIG,
)
with llm:
output = llm.generate(["Hello"], SamplingParams(max_tokens=4))
assert len(output) == 1
Expand All @@ -193,14 +213,17 @@ def test_e2e_text_26b_dummy():


@requires_gemma4_transformers
@pytest.mark.skipif(not _model_available("E2B"), reason="gemma-4-E2B-it not found")
def test_e2e_text_e2b_dummy():
"""E2E text generation for E2B (KV sharing + PLE + double-wide MLP)."""
from tensorrt_llm.llmapi import LLM, SamplingParams

dummy_dir = _make_dummy_config_dir(MODEL_PATHS["E2B"])
try:
llm = LLM(dummy_dir, load_format="dummy", attn_backend="FLASHINFER", dtype="bfloat16")
llm = LLM(
dummy_dir,
load_format="dummy",
attn_backend="FLASHINFER",
dtype="bfloat16",
kv_cache_config=_KV_CACHE_CONFIG,
)
with llm:
output = llm.generate(["Hello"], SamplingParams(max_tokens=4))
assert len(output) == 1
Expand All @@ -210,14 +233,57 @@ def test_e2e_text_e2b_dummy():


@requires_gemma4_transformers
@pytest.mark.skipif(not _model_available("31B"), reason="gemma-4-31B-it not found")
def test_e2e_text_31b_dummy():
"""E2E text generation for 31B (K=V + hybrid attn + softcap)."""
from tensorrt_llm.llmapi import LLM, SamplingParams

dummy_dir = _make_dummy_config_dir(MODEL_PATHS["31B"])
try:
llm = LLM(dummy_dir, load_format="dummy", attn_backend="FLASHINFER", dtype="bfloat16")
llm = LLM(
dummy_dir,
load_format="dummy",
attn_backend="FLASHINFER",
dtype="bfloat16",
kv_cache_config=_KV_CACHE_CONFIG,
)
with llm:
output = llm.generate(["Hello"], SamplingParams(max_tokens=4))
assert len(output) == 1
assert len(output[0].outputs[0].token_ids) > 0
finally:
shutil.rmtree(dummy_dir, ignore_errors=True)


@requires_gemma4_transformers
@pytest.mark.parametrize("max_seq_len", [256, 512])
def test_e2e_text_e2b_dummy_small_max_seq_len(max_seq_len):
"""E2E with max_seq_len <= sliding_window (256 and 512).

Regression test: when max_seq_len is at most the sliding window size, all
attention windows clamp to max_seq_len, the KV cache manager reports a
single window (is_vswa False), and the FlashInfer metadata used to skip
the per-pool page-index mapping — while the pools still differ by
head_dim. The shared page-index list then sends out-of-range page ids to
the smaller (global head_dim) pool and append_paged_kv_cache crashes with
an illegal memory access during the warmup prefill.
"""
# Real E2B geometry with dummy weights: the sliding (head_dim 256) and
# global (head_dim 512) layers must keep different page-index scales,
# and sliding_window must stay 512 so both max_seq_len values are at
# most the window. Shrinking the config changes both and hides the bug.
dummy_dir = _make_dummy_config_dir(
MODEL_PATHS["E2B"],
dummy_head_dim=256,
dummy_global_head_dim=512,
shrink_hidden=False,
)
try:
llm = LLM(
dummy_dir,
load_format="dummy",
attn_backend="FLASHINFER",
dtype="bfloat16",
max_seq_len=max_seq_len,
kv_cache_config=_KV_CACHE_CONFIG,
)
with llm:
output = llm.generate(["Hello"], SamplingParams(max_tokens=4))
assert len(output) == 1
Expand All @@ -227,14 +293,17 @@ def test_e2e_text_31b_dummy():


@requires_gemma4_transformers
@pytest.mark.skipif(not _model_available("E4B"), reason="gemma-4-E4B-it not found")
def test_e2e_text_e4b_dummy():
"""E2E text generation for E4B (KV sharing + hybrid attn)."""
from tensorrt_llm.llmapi import LLM, SamplingParams

dummy_dir = _make_dummy_config_dir(MODEL_PATHS["E4B"])
try:
llm = LLM(dummy_dir, load_format="dummy", attn_backend="FLASHINFER", dtype="bfloat16")
llm = LLM(
dummy_dir,
load_format="dummy",
attn_backend="FLASHINFER",
dtype="bfloat16",
kv_cache_config=_KV_CACHE_CONFIG,
)
with llm:
output = llm.generate(["Hello"], SamplingParams(max_tokens=4))
assert len(output) == 1
Expand All @@ -249,14 +318,11 @@ def test_e2e_text_e4b_dummy():


@requires_gemma4_transformers
@pytest.mark.skipif(not _model_available("26B"), reason="gemma-4-26B-A4B-it not found")
def test_e2e_multimodal_26b_dummy():
"""E2E multimodal: image → vision tower → embedder → LLM → output."""
import numpy as np
from PIL import Image

from tensorrt_llm.llmapi import LLM, SamplingParams

dummy_dir = _make_dummy_config_dir(MODEL_PATHS["26B"])
try:
# Format prompt with image placeholder via tokenizer chat template
Expand All @@ -277,7 +343,13 @@ def test_e2e_multimodal_26b_dummy():
tokenize=False,
)

llm = LLM(dummy_dir, load_format="dummy", attn_backend="FLASHINFER", dtype="bfloat16")
llm = LLM(
dummy_dir,
load_format="dummy",
attn_backend="FLASHINFER",
dtype="bfloat16",
kv_cache_config=_KV_CACHE_CONFIG,
)
with llm:
img = Image.fromarray(np.random.randint(0, 255, (224, 224, 3), dtype=np.uint8))
prompt = {
Expand Down
Loading