diff --git a/docs/source/features/kvcache.md b/docs/source/features/kvcache.md index cdbc908477c7..ffbf04917dc6 100644 --- a/docs/source/features/kvcache.md +++ b/docs/source/features/kvcache.md @@ -89,6 +89,7 @@ Models that select the V2 manager by default: | DeepSeek-V4 | Sparse attention attaches auxiliary per-layer buffers | | GPT-OSS | Sliding window on every other layer (VSWA), so the sliding-window and full-attention pools are sized independently | | Gemma3 / Gemma4 (text and multimodal) | Alternating sliding-window and full-attention layers (VSWA); same independent pool sizing | +| Llama / Llama4 | Uniform KV pool layout (chunked attention does not partition the pools); validated across text, multimodal, and disaggregated workloads | Separately, Gemma4 hybrid attention and sparse-attention models are routed to V2 unconditionally: their per-layer buffer layouts cannot be represented by V1's diff --git a/tensorrt_llm/_torch/models/modeling_llama.py b/tensorrt_llm/_torch/models/modeling_llama.py index 9edf8163f0ef..4c4b42241164 100644 --- a/tensorrt_llm/_torch/models/modeling_llama.py +++ b/tensorrt_llm/_torch/models/modeling_llama.py @@ -1134,6 +1134,14 @@ def forward( @register_auto_model("LlamaForCausalLM") class LlamaForCausalLM(SpecDecOneEngineForCausalLM[LlamaModel, LlamaConfig]): + @classmethod + def get_preferred_kv_cache_manager_version( + cls, + pretrained_config: Any = None, + ) -> Literal["V2"]: + """Prefer KV cache manager V2 for Llama.""" + return "V2" + @classmethod def get_preferred_transceiver_runtime( cls, @@ -1505,6 +1513,22 @@ def call_with_text_prompt( class Llama4ForConditionalGeneration(SpecDecOneEngineForCausalLM[Llama4Model, Llama4Config]): + @classmethod + def get_preferred_kv_cache_manager_version( + cls, + pretrained_config: Any = None, + ) -> Literal["V2"]: + """Prefer KV cache manager V2 for Llama4.""" + return "V2" + + @classmethod + def get_preferred_transceiver_runtime( + cls, + pretrained_config: Any = None, + ) -> Optional[Literal["CPP", "PYTHON"]]: + """Prefer the Python transceiver for Llama4 NIXL disaggregated serving.""" + return "PYTHON" + def __init__( self, model_config: ModelConfig[Llama4Config], diff --git a/tests/integration/defs/kv_cache/test_kv_cache_iteration_stats.py b/tests/integration/defs/kv_cache/test_kv_cache_iteration_stats.py index 6319aa85976e..5b4c6dd7fbfd 100644 --- a/tests/integration/defs/kv_cache/test_kv_cache_iteration_stats.py +++ b/tests/integration/defs/kv_cache/test_kv_cache_iteration_stats.py @@ -49,10 +49,18 @@ "primaryMaxNumBlocks", "primaryFreeNumBlocks", "primaryUsedNumBlocks", + "primaryEvictableNumBlocks", + "primaryPeakFreeNumBlocks", + "primaryPeakUsedNumBlocks", + "primaryPeakEvictableNumBlocks", # Instantaneous gauges — secondary (host) pool "secondaryMaxNumBlocks", "secondaryFreeNumBlocks", "secondaryUsedNumBlocks", + "secondaryEvictableNumBlocks", + "secondaryPeakFreeNumBlocks", + "secondaryPeakUsedNumBlocks", + "secondaryPeakEvictableNumBlocks", # Per-iteration deltas — context phase "iterAllocTotalBlocks", "iterAllocNewBlocks", @@ -71,8 +79,22 @@ # Intra-device (GPU → GPU) block copies "iterIntraDeviceCopyBlocks", "iterIntraDeviceCopyBytes", + # Host blocks dropped instead of being copied to another cold tier + "iterHostDroppedBlocks", + "iterHostDroppedBytes", ] +SECONDARY_FIELDS = { + "secondaryMaxNumBlocks", + "secondaryFreeNumBlocks", + "secondaryUsedNumBlocks", + "secondaryEvictableNumBlocks", + "secondaryPeakFreeNumBlocks", + "secondaryPeakUsedNumBlocks", + "secondaryPeakEvictableNumBlocks", +} +NON_SECONDARY_FIELDS = set(ALL_FIELDS) - SECONDARY_FIELDS + TEST_NAMES = { 1: "Cold start", 2: "Partial block reuse", @@ -383,26 +405,36 @@ def test_rapid_fire(self, llm_instance, all_collected, request): assert total_alloc > 0, "iterAllocTotalBlocks = 0 across all entries" def test_field_completeness(self, llm_instance, all_collected, request): - """Field completeness — verify all 18 fields present across all collected stats.""" + """Field completeness — verify fields in their manager-specific views.""" # If running standalone (no prior tests), generate some traffic if not all_collected: llm_instance.generate(["Hello world"], SamplingParams(max_tokens=16)) collect_stats(llm_instance, all_collected) entries_with_kv = 0 - missing_fields = set() for s in all_collected: ki = s.get("kvCacheIterationStats") if ki: entries_with_kv += 1 - for ws, v in ki.items(): - for field in ALL_FIELDS: - if field not in v: - missing_fields.add(field) + is_v2 = "kvCacheIterationStatsByPoolGroup" in s + expected_fields = NON_SECONDARY_FIELDS if is_v2 else set(ALL_FIELDS) + for v in ki.values(): + missing_fields = expected_fields - v.keys() + assert not missing_fields, ( + f"Missing kvCacheIterationStats fields: {sorted(missing_fields)}" + ) + + if is_v2 and "kvCacheIterationStatsByColdPoolGroup" in s: + cold_stats = s["kvCacheIterationStatsByColdPoolGroup"] + for v in cold_stats.values(): + missing_fields = SECONDARY_FIELDS - v.keys() + assert not missing_fields, ( + "Missing kvCacheIterationStatsByColdPoolGroup fields: " + f"{sorted(missing_fields)}" + ) print(f" Entries with kvCacheIterationStats: {entries_with_kv}/{len(all_collected)}") assert entries_with_kv > 0, "no entries contain kvCacheIterationStats" - assert len(missing_fields) == 0, f"Missing fields: {sorted(missing_fields)}" # --------------------------------------------------------------------------- diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index b4b1569ed66d..fbfa32bffe85 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -711,6 +711,8 @@ def test_registered_models_prefer_v2(self): "Gemma4ForCausalLM", "Gemma4ForConditionalGeneration", "Gemma4UnifiedForConditionalGeneration", + "LlamaForCausalLM", + "Llama4ForConditionalGeneration", ) for architecture in architectures: model_cls = get_registered_model_class(architecture) @@ -751,6 +753,8 @@ def test_registered_models_keep_v2_on_nixl(self): "Gemma4ForCausalLM", "Gemma4ForConditionalGeneration", "Gemma4UnifiedForConditionalGeneration", + "LlamaForCausalLM", + "Llama4ForConditionalGeneration", ) for architecture in architectures: model_cls = get_registered_model_class(architecture)