# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

import copy
import dataclasses
import os
from typing import Any, Dict, List, Optional, Sequence, Union

import torch

import tensorrt_llm
import tensorrt_llm.bindings.executor as trtllm
from tensorrt_llm._utils import (confidential_compute_enabled, get_sm_version,
                                 is_sm_100f, prefer_pinned,
                                 str_dtype_to_binding, torch_dtype_to_str)
from tensorrt_llm.inputs.multimodal import MultimodalParams

# isort: off
from tensorrt_llm.llmapi.llm_args import (
    CacheTransceiverConfig, CapacitySchedulerPolicy, EagleDecodingConfig,
    KVEventsConfig, KvCacheCompressionConfig, KvCacheConfig, MTPDecodingConfig,
    MultimodalEncoderSchedulingPolicy, PeftCacheConfig, SchedulerConfig,
    SparseAttentionConfig, SpeculativeConfig, TorchLlmArgs, WaitingQueuePolicy)
# isort: on
from tensorrt_llm._torch.peft.lora.config import (
    LoraConfig, get_default_trtllm_modules_to_hf_modules)
from tensorrt_llm._torch.peft.lora.manager import (load_torch_lora,
                                                   supports_native_fp8_lora)
from tensorrt_llm.logger import logger
from tensorrt_llm.mapping import CpType, Mapping
from tensorrt_llm.quantization import QuantAlgo

from ..attention.backends import get_sparse_attn_kv_cache_manager
from ..disaggregation.kv_cache_transceiver import (
    AttentionTypeCpp, create_kv_cache_transceiver,
    maybe_enable_fabric_memory_for_python_transceiver)
from ..hostfunc import set_low_latency_dispatch
from ..model_config import ModelConfig
from ..models.modeling_multimodal_mixin import MultimodalModelMixin
from ..speculative import (draft_prompt_lookahead, get_num_extra_kv_tokens,
                           get_num_spec_layers, get_spec_decoder,
                           should_use_separate_draft_kv_cache)
from ..utils import is_gdn_replay_enabled
from .config_utils import (MambaKVCacheParams, _is_sliding_attention_layer,
                           extract_mamba_kv_cache_params,
                           extract_qwen4_exp_ple_cache_params,
                           get_layer_attention_window, is_gemma4_hybrid,
                           is_hybrid_linear, is_kimi_linear, is_mla,
                           is_nemotron_hybrid, is_qwen3_hybrid, is_qwen4_exp,
                           resolve_vocab_size, uses_vswa_kv_cache_layout)
from .connectors.kv_cache_connector import KvCacheConnectorManager
from .dwdp import DwdpManager
from .guided_decoder import GuidedDecoder
from .kv_cache.kv_cache_manager_v2 import KVCacheManagerV2
from .kv_cache.mamba_cache_manager import (BaseMambaCacheManager,
                                           CppMambaHybridCacheManager,
                                           MambaHybridCacheManagerV2,
                                           MixedMambaHybridCacheManager,
                                           use_py_mamba_cache_manager)
from .llm_request import ExecutorResponse, LlmRequestState
from .model_engine import PyTorchModelEngine
from .py_executor import PyExecutor
from .resource_manager import (KVCacheCompressionManager, KVCacheManager,
                               PeftCacheManager, ResourceManager,
                               ResourceManagerType)
from .sampler import EarlyStopSampler, EarlyStopWithMMResult, TorchSampler
from .scheduler import (BindCapacityScheduler, BindMicroBatchScheduler,
                        KVCacheV2Scheduler, MultimodalEagerEncoderScheduler,
                        MultimodalScheduler, SimpleScheduler,
                        SimpleUnifiedScheduler)
from .seq_slot_manager import SeqSlotManager

GB = 1 << 30


def ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b


def _get_initial_lora_data_type(
    configured_lora_data_type: Optional[torch.dtype],
) -> Optional[torch.dtype]:
    if configured_lora_data_type != torch.float8_e4m3fn:
        return None
    if supports_native_fp8_lora(torch.cuda.get_device_capability()):
        return configured_lora_data_type
    return None


def _non_hybrid_kv_cache_manager_cls(config, kv_cache_config: KvCacheConfig):
    # Models with per-layer head_dim (e.g., Gemma4 hybrid attention)
    # require KVCacheManagerV2 for per-layer buffer sizes.
    needs_v2 = (kv_cache_config.use_kv_cache_manager_v2 is True
                or is_gemma4_hybrid(config))
    return KVCacheManagerV2 if needs_v2 else KVCacheManager


def kv_cache_manager_v2_incompatible_features(
        max_beam_width: Optional[int]) -> List[str]:
    """Runtime features a V2 manager cannot serve.

    ``KvCacheCreator._validate_or_fallback_kv_cache_manager_v2`` demotes a plain
    V2 manager to ``KVCacheManager`` when this list is non-empty, and rejects the
    model families that require V2 outright. ``resolved_kv_cache_manager_is_v2``
    reads the same list so that callers sizing pools from the manager version
    cannot disagree with the selection itself.

    A KV connector is deliberately not a trigger: it is served through the pool
    layout registration path and no longer forces a fallback. It is not taken as
    a parameter either, so a future caller cannot reintroduce the demotion by
    passing it.
    """
    incompat: List[str] = []
    if max_beam_width is not None and max_beam_width > 1:
        incompat.append("max_beam_width > 1")
    return incompat


def resolved_kv_cache_manager_is_v2(kv_cache_config: KvCacheConfig,
                                    max_beam_width: Optional[int]) -> bool:
    """Whether the executor will actually hold a V2 manager.

    ``use_kv_cache_manager_v2`` is a request, not the outcome: model loading has
    already resolved ``"auto"``, but a plain model is still demoted to V1 at
    manager-selection time when its runtime features are V2-incompatible. Sizing
    a pool from the request would leave V2 geometry on a V1 executor.
    """
    return (kv_cache_config.use_kv_cache_manager_v2 is True
            and not kv_cache_manager_v2_incompatible_features(max_beam_width))


def _resolve_disagg_transceiver_route(
    cache_transceiver_config: Optional[CacheTransceiverConfig],
) -> tuple[Optional[str], Optional[str]]:
    """Return the effective backend and runtime used for manager routing."""
    if cache_transceiver_config is None:
        return None, None

    backend, _ = cache_transceiver_config._resolve_default_backend()
    runtime = cache_transceiver_config.transceiver_runtime
    if runtime == "auto":
        # Model loading normally resolves ``auto``. Paths that skip model
        # defaults use the global C++ fallback, matching transceiver creation.
        runtime = None
    return backend, runtime


def is_disagg_enabled(
        cache_transceiver_config: Optional[CacheTransceiverConfig]) -> bool:
    """Whether this executor participates in disaggregated serving."""
    return (cache_transceiver_config is not None
            and cache_transceiver_config.backend is not None)


def get_kv_cache_manager_cls(
        model_config: ModelConfig,
        kv_cache_config: KvCacheConfig,
        is_disagg: bool = False,
        cache_transceiver_config: Optional[CacheTransceiverConfig] = None):
    """Resolve the concrete KV cache manager class for ``model_config``.

    For hybrid mamba models the choice between
    ``MambaHybridCacheManagerV2`` and compatibility managers is made here.
    Callers that don't care about disagg can omit ``is_disagg`` and get the
    unified-pool default.

    Model loading resolves ``use_kv_cache_manager_v2="auto"`` to V2 for
    supported hybrid Mamba models. An explicit ``False`` selects a
    compatibility manager. In disaggregated serving, V2 additionally requires
    the Python transceiver with the NIXL backend. Unsupported V2 routes fail
    rather than falling back to a different manager.

    Env-var overrides:
      * ``TRTLLM_USE_PY_MAMBA=1``  — Mixed manager in aggregated serving.
      * ``TLLM_MAMBA_MANAGER_PREFERENCE`` — explicit manager preference.
    """
    config = model_config.pretrained_config
    sparse_attn_config = model_config.sparse_attention_config
    sparse_attn_algorithm = getattr(sparse_attn_config, "algorithm", None)
    quant_config = getattr(model_config, "quant_config", None)
    if (sparse_attn_config is None and is_mla(config)
            and quant_config is not None
            and quant_config.quant_mode.has_fp4_kv_cache()):
        if kv_cache_config.use_kv_cache_manager_v2 is False:
            raise ValueError("FP4 MLA requires use_kv_cache_manager_v2=True.")
        if model_config.attn_backend != "TRTLLM":
            raise ValueError("FP4 MLA requires the TRTLLM attention backend.")
        if is_disagg:
            raise NotImplementedError(
                "FP4 MLA disaggregated serving requires the follow-up "
                "Python NIXL integration.")
        if is_hybrid_linear(config):
            raise NotImplementedError(
                "FP4 MLA requires Fp4MlaKVCacheManagerV2, which does not "
                "support hybrid linear-attention models.")
        from ..attention.backends.fp4_mla.cache_manager import \
            Fp4MlaKVCacheManagerV2

        return Fp4MlaKVCacheManagerV2
    use_v2 = kv_cache_config.use_kv_cache_manager_v2 is True
    if is_hybrid_linear(config):
        # Degenerate case: model is flagged as hybrid but the config has zero
        # mamba layers. Fall through to the standard non-hybrid routes.
        if model_config.get_num_mamba_layers() == 0:
            logger.info("Hybrid linear model has 0 mamba layers; using "
                        "KV cache manager without mamba caching")
            if sparse_attn_config is not None:
                return get_sparse_attn_kv_cache_manager(
                    sparse_attn_config, use_kv_cache_manager_v2=use_v2)
            return _non_hybrid_kv_cache_manager_cls(config, kv_cache_config)

        if sparse_attn_algorithm == "qsa" and not use_v2:
            raise ValueError(
                "QSA with hybrid Mamba / linear-attention models requires "
                "use_kv_cache_manager_v2=True.")
        if (sparse_attn_config is not None
                and sparse_attn_algorithm not in ("qsa", "skip_softmax")):
            raise ValueError(
                f"Sparse attention algorithm {sparse_attn_algorithm!r} is not "
                "supported with hybrid Mamba / linear-attention models.")

        state_config = kv_cache_config.mamba_state_config
        has_additional_snapshots = bool(
            state_config.additional_snapshot_offsets_from_start
            or state_config.additional_snapshot_offsets_from_end)
        if has_additional_snapshots and not use_v2:
            raise ValueError("Mamba additional snapshot offsets require "
                             "use_kv_cache_manager_v2=True; V1 supports only "
                             "periodic_snapshot_interval.")

        # Kimi K3 (KDA + MLA hybrid): block reuse uses the unified C++ pool
        # (CppMambaHybridCacheManager) like the other hybrid linear models —
        # per-block KDA state snapshots every mamba_state_cache_interval
        # tokens with FORCE_CHUNK context chunking. Without block reuse the
        # Mixed manager (separate KV / recurrent-state pools) stays the
        # default. SA speculative decoding is validated on the Mixed
        # manager's SpeculativeState scratch path only; reuse + SA is
        # unvalidated. Disaggregated serving (TRTLLM-14815) routes through
        # the shared hybrid transceiver validation below: the Python NIXL
        # transceiver selects the Mixed manager, whose KDA recurrent/conv
        # states transfer through the bounce buffer.
        # Helix x speculation bookkeeping (per-token verify groups on the
        # superblock ledger, py_helix_decode_group_index advancement) exists
        # only in KVCacheManagerV2. The V1-family hybrid managers account
        # helix decode one token per iteration and have no helix-x-spec
        # path, so a default (V1) resolution would run silently wrong.
        if (model_config.mapping is not None
                and model_config.mapping.has_cp_helix()
                and model_config.spec_config is not None and not use_v2):
            raise ValueError(
                "Helix with speculative decoding requires "
                "kv_cache_config.use_kv_cache_manager_v2=True; the V1-family "
                "hybrid managers do not implement per-token verify-group "
                "bookkeeping.")

        if is_kimi_linear(config) and not use_v2 and not is_disagg:
            if kv_cache_config.enable_block_reuse:
                logger.info(
                    "Using CppMambaHybridCacheManager for Kimi K3 hybrid "
                    "model (block reuse enabled)")
                return CppMambaHybridCacheManager
            logger.info(
                "Using MixedMambaHybridCacheManager for Kimi K3 hybrid model")
            return MixedMambaHybridCacheManager

        # Skip Softmax only changes attention kernels. Hybrid models still
        # need a Mamba-capable cache manager for recurrent state.
        if is_disagg:
            backend, runtime = _resolve_disagg_transceiver_route(
                cache_transceiver_config)
            if is_kimi_linear(config) and (runtime != "PYTHON"
                                           or backend != "NIXL"):
                # Only the Python NIXL transceiver can move KDA recurrent
                # state; the C++ transceiver would silently serve wrong
                # results. Model loading resolves ``auto`` to PYTHON via
                # KimiLinearForCausalLM.get_preferred_transceiver_runtime
                # (NIXL-gated); this rejects explicit non-Python routes and
                # paths that skip model defaults.
                raise ValueError(
                    "Kimi K3 disaggregated serving requires the Python "
                    "transceiver: set cache_transceiver_config "
                    "backend='NIXL' with transceiver_runtime='PYTHON' (or "
                    "leave transceiver_runtime='auto' with the NIXL "
                    "backend). The C++ transceiver cannot transfer KDA "
                    f"recurrent state (got backend={backend!r}, "
                    f"transceiver_runtime={runtime!r}).")
            if use_v2:
                if runtime != "PYTHON" or backend != "NIXL":
                    raise ValueError(
                        "KV cache manager V2 for hybrid Mamba disaggregated "
                        "serving requires transceiver_runtime='PYTHON' with "
                        "backend='NIXL'.")
            else:
                if (kv_cache_config.enable_block_reuse and runtime == "PYTHON"):
                    raise ValueError(
                        "Hybrid Mamba disaggregated serving with block reuse "
                        "and transceiver_runtime='PYTHON' requires "
                        "use_kv_cache_manager_v2=True.")
                if kv_cache_config.enable_block_reuse:
                    return CppMambaHybridCacheManager
                if runtime == "PYTHON" and backend == "NIXL":
                    logger.info("Python transceiver detected; using "
                                "MixedMambaHybridCacheManager for hybrid model")
                    return MixedMambaHybridCacheManager
                return CppMambaHybridCacheManager

        if use_py_mamba_cache_manager() and not is_disagg:
            if use_v2:
                raise ValueError(
                    "TRTLLM_USE_PY_MAMBA=1 conflicts with explicit "
                    "use_kv_cache_manager_v2=True.")
            if kv_cache_config.enable_block_reuse:
                raise ValueError(
                    "TRTLLM_USE_PY_MAMBA=1 forces "
                    "MixedMambaHybridCacheManager, which does not support "
                    "block reuse. Disable block reuse or unset "
                    "TRTLLM_USE_PY_MAMBA to use the configured cache manager.")
            logger.info(
                "Using MixedMambaHybridCacheManager for hybrid mamba model")
            return MixedMambaHybridCacheManager
        env_override = os.environ.get('TLLM_MAMBA_MANAGER_PREFERENCE', None)
        if env_override is not None:
            env_override = env_override.upper()
            if env_override == 'MIXED':
                if use_v2:
                    raise ValueError(
                        "TLLM_MAMBA_MANAGER_PREFERENCE=MIXED conflicts with "
                        "explicit use_kv_cache_manager_v2=True.")
                if kv_cache_config.enable_block_reuse:
                    raise ValueError(
                        "TLLM_MAMBA_MANAGER_PREFERENCE=MIXED forces "
                        "MixedMambaHybridCacheManager, which does not support "
                        "block reuse. Disable block reuse, use the CPP "
                        "preference, or explicitly enable KV cache manager "
                        "V2.")
                logger.warning(
                    "Environment variable TLLM_MAMBA_MANAGER_PREFERENCE=MIXED "
                    "overrides the default Mamba cache manager to "
                    "MixedMambaHybridCacheManager.")
                return MixedMambaHybridCacheManager
            if env_override == 'CPP':
                if use_v2:
                    raise ValueError(
                        "TLLM_MAMBA_MANAGER_PREFERENCE=CPP conflicts with "
                        "explicit use_kv_cache_manager_v2=True.")
                logger.warning(
                    "Environment variable TLLM_MAMBA_MANAGER_PREFERENCE=CPP "
                    "overrides the default Mamba cache manager to "
                    "CppMambaHybridCacheManager.")
                return CppMambaHybridCacheManager
            logger.warning(
                f"Unrecognized value for TLLM_MAMBA_MANAGER_PREFERENCE: {env_override}. "
                "Expected 'CPP' or 'MIXED'. Using the configured "
                "KV cache manager default.")

        if not use_v2:
            return CppMambaHybridCacheManager

        if (kv_cache_config.enable_block_reuse
                and kv_cache_config.enable_kv_pool_rebalance):
            raise ValueError(
                "V2 Mamba block reuse is not compatible with "
                "enable_kv_pool_rebalance because the rebalancer does not "
                "yet model retained recurrent-state snapshots.")
        if sparse_attn_algorithm == "qsa":
            return get_sparse_attn_kv_cache_manager(
                sparse_attn_config, use_kv_cache_manager_v2=True)
        return MambaHybridCacheManagerV2
    elif sparse_attn_config is not None:
        if sparse_attn_algorithm == "qsa":
            # The QSA manager extends the hybrid manager because it must retain
            # both paged attention data and recurrent GDN state.
            raise ValueError(
                "QSA sparse attention currently requires a hybrid Mamba / "
                "linear-attention model.")
        return get_sparse_attn_kv_cache_manager(sparse_attn_config,
                                                use_kv_cache_manager_v2=use_v2)
    else:
        return _non_hybrid_kv_cache_manager_cls(config, kv_cache_config)


# --- KV cache cost model ------------------------------------------------------
#
# KVCacheManager.get_cache_size_per_token may return either an ``int``
# (legacy proportional model ``bytes = slope * tokens``) or an affine
# ``(slope, intercept)`` tuple (CppMambaHybridCacheManager, where mamba
# state introduces a per-batch fixed cost). CacheCost normalizes the combined
# shape so the rest of the file does plain attribute access and method calls
# instead of branching on type. Managers are responsible for including any
# pool-specific alignment or capacity headroom in the returned cost.


@dataclasses.dataclass(frozen=True)
class CacheCost:
    """Affine KV cache budget: ``bytes = slope * tokens + intercept``.

    The legacy proportional case is just ``intercept = 0``.
    """
    slope: int
    intercept: int = 0

    @classmethod
    def from_raw(cls, raw) -> "CacheCost":
        """Wrap an int / tuple / CacheCost result uniformly."""
        if isinstance(raw, CacheCost):
            return raw
        if isinstance(raw, tuple):
            slope, intercept = raw
            return cls(slope=int(slope), intercept=int(intercept))
        return cls(slope=int(raw))

    def __add__(self, other: "CacheCost") -> "CacheCost":
        if not isinstance(other, CacheCost):
            return NotImplemented
        return CacheCost(slope=self.slope + other.slope,
                         intercept=self.intercept + other.intercept)

    def __str__(self) -> str:
        if self.intercept == 0:
            return f"{self.slope} bytes/token"
        else:
            return f"{self.slope} bytes/token + {self.intercept} bytes fixed cost"

    def tokens_for_budget(self, budget: int) -> int:
        """Memory budget -> max tokens. Clamps a negative result to 0."""
        if self.slope <= 0:
            return 0
        tokens = max((budget - self.intercept) // self.slope, 0)
        return tokens

    def bytes_for_tokens(self, tokens: int) -> int:
        """Token count -> memory bytes."""
        return self.slope * tokens + self.intercept


def get_attention_workspace_bytes_per_token(model_config, mapping) -> int:
    """Per-token workspace headroom the model's selected attention backend declares.

    The KV-cache profiling forward under-measures any attention workspace sized by a runtime quantity it
    does not drive to its serving maximum (e.g. ``total_kv_len``, inflated by KV reuse). Backends declare
    such a buffer via ``AttentionBackend.runtime_workspace_bytes_per_token``; this resolves the model's
    backend and returns its rate. A backend that stages no such buffer inherits the default 0, so no
    workspace is reserved and no admission cap is installed for it. See ``ATTENTION_DEVELOPER_GUIDE.md``
    §2.3.
    """
    from ..attention.backends.utils import get_attention_backend

    # Resolved without ``sparse_params``: those are per-layer, while this workspace is one buffer shared
    # across every attention layer, so the declaration is a whole-model question. The sparse backends all
    # derive from the dense class and inherit its declaration, which reads the sparse gate off model_config.
    return get_attention_backend(
        model_config.attn_backend).runtime_workspace_bytes_per_token(
            model_config, mapping)


def get_attention_workspace_is_chunked_prefill_bounded(model_config) -> bool:
    """Whether chunked prefill bounds the selected backend's runtime workspace."""
    from ..attention.backends.utils import get_attention_backend

    return get_attention_backend(
        model_config.attn_backend).runtime_workspace_is_chunked_prefill_bounded(
            model_config)


def get_mla_context_workspace_kv_len_cap(
        kv_cache_config,
        max_batch_size,
        max_num_tokens,
        max_seq_len,
        enable_chunked_prefill,
        workspace_is_chunked_prefill_bounded=True,
        chunked_workspace_profiled=True,
        require_chunked_workspace_profile=True):
    """Max summed attended-KV length covered by the context-MLA workspace reserve.

    KV-cache reuse can grow this workspace beyond the fresh-prefill profiling
    floor. Chunked prefill normally prevents that by staging one bounded KV
    chunk per launch. An implementation that consumes the complete attended
    prefix sets ``workspace_is_chunked_prefill_bounded=False`` and receives the
    same reservation and scheduler admission protection as cache reuse. A
    bounded chunk must also be exercised by profiling before dropping its
    reservation (``chunked_workspace_profiled``) when
    ``require_chunked_workspace_profile`` is enabled. Older hardware keeps
    the existing backend-declared reserve policy without this new requirement.

    Otherwise the default (no override) is the never-stall worst case ``min(max_batch_size, max_num_tokens)
    * max_seq_len``: at most that many context requests run in a step, each attending at most ``max_seq_len``
    KV, so reserving for it never defers a request. An explicit ``fp8_context_mla_kv_len_cap`` override
    reserves less workspace (freeing KV pool) and lets the scheduler defer over-cap requests; it is floored
    at ``max_seq_len`` (one request must always fit) and capped at the worst case.
    """
    workspace_can_exceed_profile = (
        kv_cache_config.enable_block_reuse and not enable_chunked_prefill) or (
            enable_chunked_prefill and
            (not workspace_is_chunked_prefill_bounded or
             (require_chunked_workspace_profile
              and not chunked_workspace_profiled)))
    if not workspace_can_exceed_profile:
        return None
    worst_case = min(max_batch_size, max_num_tokens) * max_seq_len
    override = kv_cache_config.fp8_context_mla_kv_len_cap
    if override is None:
        return worst_case
    return min(max(int(override), max_seq_len), worst_case)


def get_mla_context_workspace_reserve(budget_bytes, k_bytes_per_token,
                                      w_bytes_per_token, kv_len_cap):
    """Bytes to reserve for the fp8 context-MLA workspace, and the token admission cap that reserve covers.

    Reserve ``w * kv_len_cap`` (the worst-case summed attended KV), clamped to the per-token split
    ``budget * w / (k + w)`` so a memory-constrained node shares the budget at a common token count rather
    than starving the KV pool. The admission cap is ``reserve / w == min(kv_len_cap, budget / (k + w))``
    tokens; the scheduler admits at most that much summed attended KV, so the fp8 dequant staging buffer
    this reserve covers stays within it. This accounts for the fp8 staging term only -- the separate BF16
    full-gather buffers on the reuse path are not yet charged here (tracked as a follow-up), so this bounds
    but does not by itself guarantee the reuse-path peak. Returns ``(reserve_bytes, cap_tokens)``, or
    ``(0, None)`` when any input is non-positive.
    """
    if not (budget_bytes > 0 and k_bytes_per_token > 0 and w_bytes_per_token > 0
            and kv_len_cap and kv_len_cap > 0):
        return 0, None
    reserve = min(
        w_bytes_per_token * kv_len_cap, budget_bytes * w_bytes_per_token /
        (k_bytes_per_token + w_bytes_per_token))
    return reserve, int(reserve / w_bytes_per_token)


def _normalize_attention_windows(
    max_attention_window: List[Optional[int]],
    max_seq_len: int,
) -> Optional[List[int]]:
    normalized = [
        max_seq_len if window is None else min(max_seq_len, window)
        for window in max_attention_window
    ]
    if all(window == max_seq_len for window in normalized):
        return None
    if len(set(normalized)) == 1:
        return [normalized[0]]
    return normalized


def _derive_layer_type_attention_windows(
    config: object,
    max_seq_len: int,
) -> Optional[List[int]]:
    """Per-layer attention windows for a mixed sliding/full-attention config.

    Returns one window per decoder layer, in global model layer order
    (sliding layers get the model's ``sliding_window``, full-attention layers
    ``max_seq_len``), or ``None`` when the schedule is uniform and the
    single-window default is already right. Layer types that
    ``get_layer_attention_window`` does not recognize as sliding are treated
    as full attention.
    """
    num_layers = getattr(config, "num_hidden_layers", None)
    if not isinstance(num_layers, int) or num_layers <= 0:
        return None
    if not getattr(config, "layer_types", None):
        return None
    try:
        windows = [
            get_layer_attention_window(config, layer_idx)
            for layer_idx in range(num_layers)
        ]
    except (NotImplementedError, ValueError) as error:
        logger.warning(
            "Unable to derive per-layer attention windows from layer_types "
            f"({error}); falling back to the single-window default.")
        return None
    resolved = _normalize_attention_windows(
        [max_seq_len if window is None else window for window in windows],
        max_seq_len)
    if resolved is None or not uses_vswa_kv_cache_layout(resolved):
        # A schedule with <2 distinct positive windows can't select a VSWA
        # (multi-pool) layout, so a derived vector would only reshape the single
        # shared pool. Leave it to the single-window default.
        return None
    return resolved


def _derive_v2_layer_type_attention_windows(
    kv_cache_config: KvCacheConfig,
    kv_cache_manager_cls: type,
    model_config: ModelConfig,
    max_seq_len: int,
) -> Optional[List[int]]:
    """Per-layer windows a KVCacheManagerV2 for `model_config` is built with.

    Shared by the static per-token cost model
    (`KvCacheCreator._per_manager_cache_cost`) and `_create_kv_cache_manager`,
    so the budget split sizes a manager from the same windows the manager
    receives. Returns `None`, leaving `kv_cache_config` as is, when the user
    supplied `max_attention_window`, for any class other than KVCacheManagerV2
    (KVCacheManager keeps the single-window default), when
    `_derive_layer_type_attention_windows` has nothing to derive, for hybrid
    linear-attention configs (their attention layers are interleaved with
    recurrent layers, and the static cost model indexes a window list by
    attention-layer position while the manager indexes it by global layer id,
    so the two would read a per-layer list differently), and when a
    user-supplied `pool_ratio` does not carry one entry per derived layer
    group (the manager would reject that arity at construction; a warning
    names the fix and the configuration keeps the single pool it was written
    for).
    """
    if kv_cache_config.max_attention_window is not None:
        return None
    if not (isinstance(kv_cache_manager_cls, type)
            and issubclass(kv_cache_manager_cls, KVCacheManagerV2)):
        return None
    config = model_config.pretrained_config
    derived_windows = _derive_layer_type_attention_windows(config, max_seq_len)
    if derived_windows is None:
        return None
    if is_hybrid_linear(config):
        return None
    pool_ratio = kv_cache_config.pool_ratio
    num_layer_groups = len(set(derived_windows))
    if pool_ratio is not None and len(pool_ratio) != num_layer_groups:
        logger.warning_once(
            f"kv_cache_config.pool_ratio has {len(pool_ratio)} entries, but the "
            "per-layer attention windows derived from layer_types form "
            f"{num_layer_groups} layer groups; keeping the single-window "
            "default. Provide one pool_ratio entry per layer group to size the "
            "sliding-window and full-attention pools separately, or set "
            "max_attention_window explicitly.",
            key="derived_attention_windows_pool_ratio_arity")
        return None
    return derived_windows


def _get_num_pool_groups_for_estimation(
    model_config: object,
    max_seq_len: int,
    fallback_attention_windows: Optional[List[Optional[int]]],
) -> int:
    """Estimate V2 pool groups from KV storage windows and layer metadata.

    Model sliding attention alone does not imply separate storage pools.
    Preserve distinctions required by hybrid layer types and page sizes.
    """
    if (fallback_attention_windows is not None
            and not is_hybrid_linear(model_config)):
        normalized_windows = _normalize_attention_windows(
            fallback_attention_windows, max_seq_len)
        return 1 if normalized_windows is None else len(set(normalized_windows))

    layer_types = getattr(model_config, "layer_types", None)
    attention_windows = None
    if (fallback_attention_windows is None and is_gemma4_hybrid(model_config)):
        attention_windows = _derive_layer_type_attention_windows(
            model_config, max_seq_len)
    # Check whether KV storage uses configured or inferred attention windows,
    # which can be independently configured from the model's use of sliding-window
    # attention for computation.
    kv_cache_uses_attention_windows = (fallback_attention_windows is not None
                                       or attention_windows is not None)

    if attention_windows is not None:
        normalized_windows = _normalize_attention_windows(
            attention_windows, max_seq_len)
        if normalized_windows is None:
            return 1
        return len(set(normalized_windows))

    if isinstance(layer_types, (list, tuple)):
        # These tags come from HF pretrained_config.layer_types.
        # As above, sliding-window attention in the model does not imply windowed
        # KV storage. Without windowed storage, sliding/full attention share a pool
        # type unless their page sizes differ (Gemma4).
        pool_types = set()
        for layer_type in layer_types:
            if (not kv_cache_uses_attention_windows
                    and not is_gemma4_hybrid(model_config) and
                (_is_sliding_attention_layer(layer_type)
                 or getattr(layer_type, "name",
                            str(layer_type)).lower() == "full_attention")):
                layer_type = "full_attention"
            pool_types.add(layer_type)
        if len(pool_types) > 1:
            return len(pool_types)

    if fallback_attention_windows is not None:
        normalized_windows = _normalize_attention_windows(
            fallback_attention_windows, max_seq_len)
        if normalized_windows is not None:
            return len(set(normalized_windows))

    return 1


def draft_config_defines_attention_layout(
    draft_pretrained_config: object, ) -> bool:
    """Return whether the draft HF config explicitly defines its attention layout.

    A ``True`` result makes the draft settings authoritative, including an
    explicit full-attention layout. For example, a config with
    ``use_sliding_window=False`` and ``sliding_window=4096`` returns ``True``:
    its layers should attend to ``max_seq_len`` instead of inheriting the
    target model's window. A config that provides none of
    ``use_sliding_window``, ``sliding_window``, or ``layer_types`` returns
    ``False`` so the legacy uniform-target fallback can be used.
    """
    return (
        getattr(draft_pretrained_config, "use_sliding_window", None) is not None
        or getattr(draft_pretrained_config, "sliding_window", None) is not None
        or bool(getattr(draft_pretrained_config, "layer_types", None)))


def _expand_attention_window_pattern_to_global_layers(
    max_attention_window: Optional[Sequence[int]],
    layer_mask: Sequence[bool],
) -> Optional[List[int]]:
    """Expand an enabled-layer pattern into physical global-layer order."""
    if max_attention_window is None:
        return None

    pattern = list(max_attention_window)
    global_windows = [pattern[0]] * len(layer_mask)
    enabled_layer_offset = 0
    for layer_idx, enabled in enumerate(layer_mask):
        if enabled:
            global_windows[layer_idx] = pattern[enabled_layer_offset %
                                                len(pattern)]
            enabled_layer_offset += 1
    return global_windows


def _derive_draft_max_attention_window(
    kv_cache_config: KvCacheConfig,
    draft_pretrained_config: object,
    max_seq_len: int,
    num_draft_layers: int,
) -> Optional[List[int]]:
    layer_windows = [
        get_layer_attention_window(draft_pretrained_config, layer_idx)
        for layer_idx in range(num_draft_layers)
    ]
    if draft_config_defines_attention_layout(draft_pretrained_config):
        draft_windows = [
            max_seq_len if window is None else window
            for window in layer_windows
        ]
        return _normalize_attention_windows(draft_windows, max_seq_len)

    if not uses_vswa_kv_cache_layout(kv_cache_config.max_attention_window):
        max_attention_window = kv_cache_config.max_attention_window
        if max_attention_window is None:
            return None
        return _normalize_attention_windows(max_attention_window, max_seq_len)

    return None


class KvCacheCreator:
    """Groups together logic related to KV cache construction."""

    # Byte budgets that back an offload tier reserved in full at manager
    # construction: the host tier is prefaulted and page-locked, the disk tier
    # is preallocated. Every live manager reserves its own, so these budgets
    # must be divided rather than handed out whole.
    _OFFLOAD_TIER_BUDGET_ATTRS = ("host_cache_size", "disk_cache_size")

    # Paired-reuse protocol flag. ``build_managers`` resolves it once and hands
    # it to both constructors; the managers must not re-derive it.
    _joint_kv_cache_reuse = False

    def __init__(
        self,
        *,
        model_engine: PyTorchModelEngine,
        draft_model_engine: Optional[PyTorchModelEngine],
        mapping: Mapping,
        net_max_seq_len: int,
        kv_connector_manager: Optional[KvCacheConnectorManager],
        max_num_tokens: int,
        max_beam_width: int,
        tokens_per_block: int,
        max_seq_len: int,
        max_batch_size: int,
        kv_cache_config: KvCacheConfig,
        llm_args: TorchLlmArgs,
        speculative_config: SpeculativeConfig,
        sparse_attention_config: SparseAttentionConfig,
        profiling_stage_data: Optional[dict],
        is_disagg: bool,
        execution_stream: Optional[torch.cuda.Stream] = None,
        draft_config: Optional[ModelConfig] = None,
        skip_est: bool = False,
    ):
        self._model_engine = model_engine
        self._draft_model_engine = draft_model_engine
        self._mapping = mapping
        self._kv_cache_config = kv_cache_config
        self._max_kv_tokens_in = self._kv_cache_config.max_tokens
        self._max_gpu_total_bytes_in = self._kv_cache_config.max_gpu_total_bytes
        self._pool_ratio_in = self._kv_cache_config.pool_ratio
        self._avg_seq_len_in = self._kv_cache_config.avg_seq_len
        self._max_num_tokens = max_num_tokens
        self._max_beam_width = max_beam_width
        self._kv_connector_manager = kv_connector_manager
        self._llm_args = llm_args
        self._speculative_config = speculative_config
        self._sparse_attention_config = sparse_attention_config
        self._tokens_per_block = tokens_per_block
        self._max_seq_len = max_seq_len
        self._max_batch_size = max_batch_size
        self._net_max_seq_len = net_max_seq_len
        self._dummy_reqs = None
        self._mla_chunked_profile_length: int | None = None
        self._dummy_encoder_inputs: List[MultimodalParams] = []
        self._profiling_stage_data = profiling_stage_data
        self._is_disagg = is_disagg
        self._cache_transceiver_config = llm_args.cache_transceiver_config
        self._execution_stream = execution_stream
        self._kv_cache_manager_cls = self._get_model_kv_cache_manager_cls(
            model_engine)
        self._is_kv_cache_manager_v2 = issubclass(self._kv_cache_manager_cls,
                                                  KVCacheManagerV2)
        self._disable_overlap_scheduler = llm_args.disable_overlap_scheduler
        self._draft_config = draft_config
        self._skip_est = skip_est
        # Admission cap (tokens of summed context attended-KV) that the fp8 context-MLA workspace reservation
        # covers, computed in configure_kv_cache_capacity and carried to the KV manager so the scheduler
        # reads it directly instead of re-deriving it from pool layout. None until reserved (or w == 0).
        self._fp8_ctx_mla_kv_len_cap = None
        self._maybe_enable_fabric_memory_for_python_transceiver()

    def _maybe_enable_fabric_memory_for_python_transceiver(self) -> None:
        maybe_enable_fabric_memory_for_python_transceiver(
            self._cache_transceiver_config, self._kv_cache_manager_cls)

    def _get_model_kv_cache_manager_cls(
        self,
        model_engine: PyTorchModelEngine,
        kv_cache_config_override: Optional[KvCacheConfig] = None,
    ):
        kv_cache_config = (kv_cache_config_override if kv_cache_config_override
                           is not None else self._kv_cache_config)
        model_config = model_engine.model.model_config
        cls = get_kv_cache_manager_cls(
            model_config,
            kv_cache_config,
            is_disagg=self._is_disagg,
            cache_transceiver_config=self._cache_transceiver_config)
        cls = self._validate_or_fallback_kv_cache_manager_v2(
            cls, model_config, kv_cache_config)
        if is_hybrid_linear(model_config.pretrained_config):
            logger.info_once(
                f"Selected hybrid KV cache manager: {cls.__name__}",
                key=f"hybrid_kv_cache_manager_{cls.__name__}")
        # Compatibility managers do not support MTP block reuse. Warn at the
        # routing site so users see the concrete manager selected for the
        # incompatible combination.
        if is_hybrid_linear(model_engine.model.model_config.pretrained_config) \
                and kv_cache_config.enable_block_reuse \
                and self._speculative_config is not None:
            if not issubclass(cls, MambaHybridCacheManagerV2):
                logger.warning(
                    "Block reuse does not work with MTP for hybrid linear models "
                    f"when using non-V2 Mamba cache manager {cls.__name__}")
        return cls

    def _validate_or_fallback_kv_cache_manager_v2(
            self,
            kv_cache_manager_cls,
            model_config: ModelConfig,
            kv_cache_config: Optional[KvCacheConfig] = None):
        config = model_config.pretrained_config
        # Use ``issubclass`` rather than identity equality so V2 subclasses
        # (e.g. ``MiniMaxM3KVCacheManagerV2`` from the sparse-attention path)
        # also go through the V2-incompatible-feature gate below.
        if issubclass(kv_cache_manager_cls, KVCacheManagerV2):
            sparse_attn_config = model_config.sparse_attention_config
            incompat = kv_cache_manager_v2_incompatible_features(
                self._max_beam_width)
            if incompat:
                incompat_str = ", ".join(incompat)
                # Never silently replace a sparse V2 manager with V1. Some
                # sparse models require V2 structurally; for models such as DSA
                # that support both managers, fallback would ignore the user's
                # explicit manager selection.
                if sparse_attn_config is not None:
                    raise NotImplementedError(
                        f"KVCacheManagerV2 for sparse-attention models "
                        f"(algorithm={sparse_attn_config.algorithm!r}) is not "
                        f"supported with "
                        f"{incompat_str}. Disable the incompatible features to "
                        f"run sparse-attention models.")
                # Gemma4 hybrid uses per-layer head_dim that V1 would coerce to
                # ``max(head_dim)``, changing per-layer KV byte sizes.
                if is_gemma4_hybrid(config):
                    raise NotImplementedError(
                        f"Gemma4 hybrid attention requires KVCacheManagerV2, "
                        f"which is not yet supported with {incompat_str}. "
                        f"Disable these features to run Gemma4 hybrid models.")
                quant_config = getattr(model_config, "quant_config", None)
                if (sparse_attn_config is None and is_mla(config)
                        and quant_config is not None
                        and quant_config.quant_mode.has_fp4_kv_cache()):
                    raise NotImplementedError(
                        "FP4 MLA requires Fp4MlaKVCacheManagerV2, which is "
                        f"not yet supported with {incompat_str}. Disable these "
                        "features to run FP4 MLA.")
                if is_hybrid_linear(config):
                    raise NotImplementedError(
                        "Hybrid Mamba cache managers do not support "
                        f"{incompat_str}; CppMambaHybridCacheManager does not "
                        "provide a compatible fallback. Use max_beam_width=1 "
                        "to run hybrid linear models.")
                # Plain V2 (explicitly enabled or selected by a model preference):
                # V2 was a preference, not a structural requirement, so we can
                # safely fall back to V1.
                logger.warning(
                    "KVCacheManagerV2 is not supported with %s. "
                    "Falling back to KVCacheManager.", incompat_str)
                return KVCacheManager
        return kv_cache_manager_cls

    def _enable_kv_cache_stats(self) -> bool:
        return (self._llm_args.enable_iter_perf_stats
                or getattr(self._llm_args, "return_perf_metrics", False))

    def _per_manager_cache_cost(self,
                                manager_cls,
                                model_config,
                                kv_cache_config: Optional[KvCacheConfig] = None,
                                *,
                                is_draft: bool = False,
                                mapping=None,
                                **extra_kwargs) -> CacheCost:
        kv_cache_config = (kv_cache_config if kv_cache_config is not None else
                           self._kv_cache_config)
        if not is_draft:
            derived_windows = _derive_v2_layer_type_attention_windows(
                kv_cache_config, manager_cls, model_config, self._max_seq_len)
            if derived_windows is not None:
                # Cost the manager with the per-layer windows
                # `_create_kv_cache_manager` builds it with, so the budget
                # split between the target and draft managers matches their
                # pools. The draft manager never derives (see there).
                kv_cache_config = kv_cache_config.model_copy(
                    update={"max_attention_window": derived_windows})
        return CacheCost.from_raw(
            manager_cls.get_cache_size_per_token(
                model_config,
                mapping if mapping is not None else self._mapping,
                tokens_per_block=self._tokens_per_block,
                max_seq_len=self._max_seq_len,
                max_batch_size=self._max_batch_size,
                max_num_tokens=self._max_num_tokens,
                kv_cache_config=kv_cache_config,
                spec_config=self._speculative_config,
                is_draft=is_draft,
                **extra_kwargs))

    def _get_one_model_draft_layer_mask(self) -> List[bool]:
        """Return the same draft-only mask used by runtime construction."""
        num_draft_layers = self._get_num_draft_layers()
        if self._speculative_config.spec_dec_mode.is_external_drafter():
            return [True] * num_draft_layers
        target_num_layers = (self._model_engine.model.model_config.
                             pretrained_config.num_hidden_layers)
        return [False] * target_num_layers + [True] * num_draft_layers

    def _get_kv_size_per_token(self,
                               kv_cache_config: Optional[KvCacheConfig] = None
                               ) -> CacheCost:
        """Aggregate KV cost across target + (optional) draft as a CacheCost.

        ``max_batch_size`` and ``kv_cache_config`` are passed unconditionally;
        managers that don't need them ignore via ``**kwargs``.
        """
        kv_cache_config = (kv_cache_config if kv_cache_config is not None else
                           self._kv_cache_config)
        model_config = self._model_engine.model.model_config
        use_separate_draft_kv_cache = (
            self._should_create_separate_draft_kv_cache())
        total = self._per_manager_cache_cost(
            self._kv_cache_manager_cls,
            model_config,
            kv_cache_config,
            use_separate_draft_kv_cache=use_separate_draft_kv_cache)
        if self._is_encoder_decoder():
            total += CacheCost.from_raw(self._get_cross_kv_size_per_token())
        draft_cost = self._get_draft_cache_cost(
            kv_cache_config,
            use_separate_draft_kv_cache=use_separate_draft_kv_cache,
        )
        if draft_cost is not None:
            total += draft_cost
        return total

    def _get_draft_cache_cost(
        self,
        kv_cache_config: KvCacheConfig,
        *,
        use_separate_draft_kv_cache: bool,
    ) -> Optional[CacheCost]:
        """Return the draft manager's standalone cache cost, if it has one.

        Under helix CP the drafter is dense rather than helix-sharded, so it is
        costed with the same repurposed mapping runtime construction uses, then
        the slope is multiplied by cp_size to express it per rank-LOCAL target
        token (the target stores only every cp_size-th page per rank).
        Intercepts are per-request rank-local bytes and stay unscaled.
        """
        draft_mapping = self._mapping
        helix_cp_scale = 1
        if self._mapping.has_cp_helix():
            draft_mapping = self._mapping.repurpose_helix_cp_to_tp()
            helix_cp_scale = self._mapping.cp_size

        def scaled(cost: CacheCost) -> CacheCost:
            return CacheCost(slope=cost.slope * helix_cp_scale,
                             intercept=cost.intercept)

        if self._draft_model_engine is not None:
            draft_model_config = self._draft_model_engine.model.model_config
            draft_kv_cache_manager_cls = self._get_model_kv_cache_manager_cls(
                self._draft_model_engine, kv_cache_config)
            return scaled(
                self._per_manager_cache_cost(draft_kv_cache_manager_cls,
                                             draft_model_config,
                                             kv_cache_config,
                                             mapping=draft_mapping))
        if use_separate_draft_kv_cache:
            # One-model draft with separate KV cache layout.
            # Pass num_layers explicitly since the HF config may report a
            # different layer count than what is actually used at runtime
            # (e.g. EAGLE3: config says 1, runtime uses 4).
            # For PP, draft layers are only on the last rank (see
            # get_pp_layers), so only that rank should include draft cost.
            # _get_draft_kv_model_config(), not _get_effective_draft_config():
            # the cost charged here must be the cost of the pool that
            # _create_one_model_draft_kv_cache_manager actually allocates.
            effective_draft_config = self._get_draft_kv_model_config()
            draft_kv_cache_config = self._get_one_model_draft_kv_cache_config(
                kv_cache_config, self._max_seq_len)
            # Resolve draft manager class from draft config — may differ
            # from target (e.g. hybrid target + plain transformer draft).
            draft_kv_cache_manager_cls = get_kv_cache_manager_cls(
                effective_draft_config,
                draft_kv_cache_config,
                is_disagg=self._is_disagg,
                cache_transceiver_config=self._cache_transceiver_config)
            draft_kv_cache_manager_cls = self._validate_or_fallback_kv_cache_manager_v2(
                draft_kv_cache_manager_cls, effective_draft_config,
                draft_kv_cache_config)
            if self._speculative_config.spec_dec_mode.is_external_drafter():
                # External drafter: layers start from 0, normal PP distribution
                return scaled(
                    self._per_manager_cache_cost(draft_kv_cache_manager_cls,
                                                 effective_draft_config,
                                                 draft_kv_cache_config,
                                                 mapping=draft_mapping,
                                                 is_draft=True))
            elif self._mapping.is_last_pp_rank():
                # EAGLE3/MTP: draft layers only on last PP rank
                return scaled(
                    self._per_manager_cache_cost(
                        draft_kv_cache_manager_cls,
                        effective_draft_config,
                        draft_kv_cache_config,
                        mapping=draft_mapping,
                        num_layers=self._get_num_draft_layers(),
                        is_draft=True))
        return None

    def _cal_max_memory(self, peak_memory, total_gpu_memory, fraction,
                        allocated_bytes: int) -> int:
        """
        Calculate the max KV cache capacity.

        NOTE: `allocated_bytes` is the total KV-cache memory that must be pre-allocated during the estimation phase (for both the main and draft models) so the estimation run can complete successfully. When computing `available_kv_mem`, add this amount back in.
        """
        kv_size_per_token = self._get_kv_size_per_token()

        available_kv_mem = (total_gpu_memory - peak_memory +
                            allocated_bytes) * fraction
        logger.info(
            f"Peak memory during memory usage profiling (torch + non-torch): {peak_memory / (GB):.2f} GiB, "
            f"available KV cache memory when calculating max tokens: {available_kv_mem / (GB):.2f} GiB, "
            f"fraction is set {fraction}, kv size per token is {kv_size_per_token}. device total memory {total_gpu_memory / (GB):.2f} GiB, "
            f"temporary kv cache memory during profiling {allocated_bytes / (GB):.2f} GiB"
        )
        return int(available_kv_mem)

    def _get_mla_chunked_profile_length(self, input_seq_len: int) -> int | None:
        """Length needed to profile two full cached-KV chunks with a full query."""
        model_config = self._model_engine.model.model_config
        # Skip-softmax preserves dense MLA. DSA-style hooks can switch to
        # absorption before this request reaches a full cached-KV chunk.
        if (not is_mla(model_config.pretrained_config)
                or model_config.attn_backend != "TRTLLM"
                or getattr(model_config.sparse_attention_config, "algorithm",
                           None) not in (None, "skip_softmax")):
            return None
        # Beam search and speculative decoding use this same dense context
        # path. Their extra KV capacity is accounted for by
        # _get_token_num_for_estimation; neither needs a separate exclusion.
        features = self._model_engine.attn_runtime_features
        if not features.chunked_prefill or features.chunk_size <= 0:
            return None
        # MLA.forward_context uses full-gather on Hopper even when scheduler
        # chunking is enabled. Only the SM100+ path bounds cached-KV staging.
        if get_sm_version() < 100:
            return None

        kv_chunk_tokens = (features.chunk_size *
                           features.chunked_prefill_buffer_batch_size)
        # The previous loop's K/V tensors can remain live while the next
        # chunk is expanded. Exercise two full cached-KV chunks in one forward
        # to include this overlap, alongside a full query budget. Round the
        # prefix up to a scheduler-step boundary so the final query is full.
        cached_tokens = (ceil_div(2 * kv_chunk_tokens, self._max_num_tokens) *
                         self._max_num_tokens)
        profile_length = cached_tokens + self._max_num_tokens
        if profile_length > input_seq_len:
            # A short-context workload may fill the KV chunk through fan-out;
            # a single long request cannot cover that case. Keep its reserve.
            return None
        return profile_length

    def _create_dummy_context_requests(
            self, input_seq_len: int) -> List[trtllm.Request]:
        # Keep the LLM dummy text-only so it can always fill max_num_tokens.
        # The MM encoder is profiled independently at its own token budget.
        requests = []
        vocab_size = self._model_engine.model.model_config.pretrained_config.vocab_size
        max_num_tokens = self._max_num_tokens
        max_beam_width = self._max_beam_width

        self._mla_chunked_profile_length = self._get_mla_chunked_profile_length(
            input_seq_len)
        if self._mla_chunked_profile_length is not None:
            input_seq_len = self._mla_chunked_profile_length
            remaining_tokens = input_seq_len
            logger.info(
                "Profiling chunked MLA with a cached prefix: "
                f"prompt length {input_seq_len}, query budget {max_num_tokens}."
            )
        else:
            input_seq_len = min(max_num_tokens, input_seq_len)
            remaining_tokens = max_num_tokens
        while remaining_tokens > 0:
            input_seq_len = min(input_seq_len, remaining_tokens)
            input_tokens = torch.randint(low=0,
                                         high=vocab_size,
                                         size=(input_seq_len, )).tolist()
            request = trtllm.Request(input_tokens,
                                     max_tokens=1,
                                     streaming=False,
                                     sampling_config=trtllm.SamplingConfig(
                                         beam_width=max_beam_width, ),
                                     output_config=trtllm.OutputConfig(),
                                     end_id=-1)
            if self._model_engine.use_mrope:
                request.py_multimodal_data = {
                    "mrope_config": {
                        "mrope_position_ids":
                        torch.zeros(3, 1, input_seq_len, dtype=torch.int32),
                        "mrope_position_deltas":
                        torch.zeros(1, 1, dtype=torch.int32)
                    }
                }
            request.py_conversation_params = None
            requests.append(request)
            remaining_tokens -= input_seq_len
        if self._mapping.enable_attention_dp:
            requests = requests * self._mapping.tp_size
        return requests

    def _create_dummy_encoder_inputs(self) -> List[MultimodalParams]:
        """Build one processed MM encoder batch at its scheduling limits."""
        if not isinstance(self._model_engine.model, MultimodalModelMixin):
            return []
        if isinstance(
                self._profiling_stage_data,
                dict) and not self._profiling_stage_data.get("enable_mm_reqs"):
            return []
        if self._llm_args.disable_mm_encoder:
            return []
        # MM E/P disaggregation may remove an otherwise exposed encoder.
        if (hasattr(self._model_engine.model, "mm_encoder")
                and self._model_engine.model.mm_encoder is None):
            return []

        input_processor = self._model_engine.input_processor
        encoder_max_num_tokens = self._model_engine.encoder_max_num_tokens
        if encoder_max_num_tokens is None or encoder_max_num_tokens <= 0:
            return []

        try:
            max_tokens_per_item = input_processor.get_mm_max_tokens_per_item(
                max_num_encoder_tokens=encoder_max_num_tokens)
            for modality, num_tokens in max_tokens_per_item.items():
                if not modality:
                    raise ValueError("Multimodal modality name cannot be empty")
                if num_tokens <= 0:
                    raise ValueError(
                        "Multimodal encoder token counts must be positive; "
                        f"got {num_tokens} for {modality}")
            max_tokens_per_item = {
                modality: num_tokens
                for modality, num_tokens in max_tokens_per_item.items()
                if num_tokens <= encoder_max_num_tokens
            }
            if not max_tokens_per_item:
                return []

            modality, num_tokens_per_item = max(
                max_tokens_per_item.items(),
                key=lambda item: (item[1], item[0]),
            )
            num_items = min(
                self._model_engine.encoder_batch_size,
                encoder_max_num_tokens // num_tokens_per_item,
            )
            mm_data = input_processor.get_dummy_mm_data(
                max_num_encoder_tokens=encoder_max_num_tokens,
                mm_counts={modality: num_items},
                dtype=self._model_engine.model.dtype,
            )
        except NotImplementedError:
            logger.info("Multimodal memory profiling skipped: "
                        f"{type(input_processor).__name__} does not implement "
                        "get_dummy_mm_data().")
            return []
        if not mm_data:
            return []
        if not isinstance(mm_data, dict):
            raise ValueError(
                "get_dummy_mm_data() must return a multimodal_data "
                "dictionary")
        return [MultimodalParams(multimodal_data=mm_data)]

    def _encode_dummy_inputs(self) -> Optional[torch.Tensor]:
        """Run the full-budget MM encoder and retain request-owned output storage."""
        if not self._dummy_encoder_inputs:
            return None

        encoder_inputs = self._dummy_encoder_inputs
        try:
            with torch.inference_mode():
                for encoder_input in encoder_inputs:
                    encoder_input.to_device(
                        "multimodal_data",
                        "cuda",
                        pin_memory=prefer_pinned(),
                        target_keywords=getattr(
                            self._model_engine.model,
                            "multimodal_data_device_paths",
                            None,
                        ),
                    )
                output = self._model_engine.model.encode_multimodal_inputs(
                    encoder_inputs)
                # Runtime item state owns detached copies rather than views of
                # an encoder batch. Reproduce that allocation boundary here.
                return output.detach().clone()
        finally:
            self._dummy_encoder_inputs = []

    def _get_multimodal_encoder_memory_reserve(self,
                                               profiled_output_bytes: int = 0
                                               ) -> int:
        """Return output and cache capacity absent from the measured peak."""
        output_budget = getattr(self._model_engine,
                                "mm_encoder_output_budget_bytes", None)
        unprofiled_output_bytes = max(0, (output_budget or 0) -
                                      profiled_output_bytes)

        model = self._model_engine.model
        cache_bytes = 0
        if (isinstance(model, MultimodalModelMixin)
                and model.encoder_cache_active
                and model.model_config.multimodal_config is not None):
            cache_bytes = (
                model.model_config.multimodal_config.encoder_cache_max_bytes)
        return unprofiled_output_bytes + cache_bytes

    def _get_token_num_for_estimation(self) -> int:
        """Compute KV cache capacity required for estimate_max_kv_cache_tokens to succeed."""
        if 'cp_type' in self._mapping.cp_config:
            raise ValueError(
                "KV cache size estimation not supported with context parallelism."
            )
        # estimate_max_kv_cache_tokens submits self._dummy_reqs
        num_cache_blocks = 0
        num_extra_tokens_per_seq = 1  # account for generated tokens
        spec_cfg = self._speculative_config
        if not self._llm_args.disable_overlap_scheduler and spec_cfg is not None:
            num_extra_tokens_per_seq += spec_cfg.tokens_per_gen_step - 1

        if spec_cfg is not None:
            num_extra_tokens_per_seq += spec_cfg.tokens_per_gen_step - 1
            num_extra_tokens_per_seq += get_num_extra_kv_tokens(spec_cfg)

        if self._dummy_reqs is None:
            self._dummy_reqs = self._create_dummy_context_requests(
                max(1, self._net_max_seq_len - 1))
            self._dummy_encoder_inputs = self._create_dummy_encoder_inputs()
        for req in self._dummy_reqs:
            num_req_tokens = len(req.input_token_ids) + num_extra_tokens_per_seq
            # Requests cannot share KV cache blocks. Round up to nearest integer multiple of block size.
            num_cache_blocks += ceil_div(num_req_tokens, self._tokens_per_block)

        # With ADP enabled, _create_dummy_context_requests produces tp_size
        # copies so each rank gets work during the estimation warmup. But the
        # scheduler distributes them evenly (1 per rank), so each rank's KV
        # cache only needs capacity for its own share, not all of them.
        if self._mapping.enable_attention_dp and self._mapping.tp_size > 1:
            num_cache_blocks = (num_cache_blocks + self._mapping.tp_size -
                                1) // self._mapping.tp_size

        # Max cuda graph warmup required tokens
        max_cuda_graph_bs = min(self._model_engine.batch_size,
                                self._model_engine._max_cuda_graph_batch_size)
        # Round up the max seq len to the block size
        max_seq_len_blocks = ceil_div(self._model_engine.max_seq_len + 1,
                                      self._tokens_per_block)
        cuda_graph_warmup_block = max_seq_len_blocks + max_cuda_graph_bs - 1
        num_cache_blocks = max(cuda_graph_warmup_block, num_cache_blocks)

        # This is the minimal blocks required to run with max bs
        # If not able to allocate self._model_engine.batch_size blocks, the max batch size should be adjusted.
        num_cache_blocks = max(num_cache_blocks, self._model_engine.batch_size)

        # KVCacheManagerV2 divides the quota derived from max_tokens across its
        # pool groups. Scale the dummy workload by the inferred group count so
        # each pool can hold a max-length request. This covers both VSWA pools
        # (distinct attention windows) and hybrid recurrent/attention pools
        # (distinct layer types without sliding windows).
        num_pool_groups = 1
        if self._is_kv_cache_manager_v2:
            model_cfg = self._model_engine.model.model_config.pretrained_config
            num_pool_groups = _get_num_pool_groups_for_estimation(
                model_cfg,
                self._model_engine.max_seq_len,
                self._kv_cache_config.max_attention_window,
            )
        num_cache_blocks *= num_pool_groups

        # Dummy context requests use the configured maximum beam width. Scale
        # their block budget by the same value so the temporary KV cache used
        # during warm-up can accommodate those requests.
        num_cache_blocks *= self._max_beam_width

        max_num_tokens_for_estimation = (num_cache_blocks *
                                         self._tokens_per_block)
        # V2 capacity is controlled by max_gpu_total_bytes; max_tokens only
        # describes the dummy workload needed for estimation.
        if self._is_kv_cache_manager_v2:
            return max_num_tokens_for_estimation

        free_mem, _ = torch.cuda.mem_get_info()
        max_memory = self._kv_cache_config.free_gpu_memory_fraction * free_mem
        kv_size_per_token = self._get_kv_size_per_token()
        max_num_tokens_in_memory = (
            kv_size_per_token.tokens_for_budget(max_memory) //
            self._tokens_per_block * self._tokens_per_block)
        return min(max_num_tokens_for_estimation, max_num_tokens_in_memory)

    def try_prepare_estimation(self) -> bool:
        """Prepare for possible KV cache capacity estimation.

        This updates `kv_cache_config` and returns a boolean indicating whether KV cache
        estimation is to be performend.
        """
        if self._skip_est:
            return False

        estimating_kv_cache = True
        if 'cp_type' in self._mapping.cp_config:
            estimating_kv_cache = False
            logger.info(
                "KV cache size estimation is not supported for context parallelism, disable it."
            )
            if (self._is_kv_cache_manager_v2
                    and self._mapping.cp_config.get('cp_type') == CpType.HELIX):
                # Promote like the encoder-decoder case so build_managers
                # runs configure_kv_cache_capacity(), which sets the quota
                # KVCacheManagerV2 requires at construction (V1 stays local).
                # HELIX only: configure_kv_cache_capacity has no sizing path
                # for other CP types and would hit its assertion.
                self._skip_est = True
        model_config = self._model_engine.model.model_config
        if model_config.attn_backend == "VANILLA":
            estimating_kv_cache = False
            logger.info(
                "KV cache size estimation is not supported for Vanilla attention backend, disable it."
            )
        if getattr(model_config, "is_encoder_decoder", False):
            # The estimation dummies are text-only, and the cross-KV block
            # accounting needs an encoder length (getEncoderOutputLen throws).
            # _skip_est (not just the local flag) so build_managers runs
            # configure_kv_cache_capacity(), which KVCacheManagerV2 needs for
            # its memory quota — the TRTLLM_SKIP_KV_CACHE_ESTIMATION=1 path.
            self._skip_est = True
            estimating_kv_cache = False
            logger.info(
                "KV cache size estimation is not supported for encoder-decoder "
                "models, disable it.")

        if estimating_kv_cache:
            estimate_max_tokens = self._get_token_num_for_estimation()
            max_tokens = min(
                estimate_max_tokens, self._kv_cache_config.max_tokens
            ) if self._kv_cache_config.max_tokens is not None else estimate_max_tokens
            # User-provided pool sizing can underprovision the temporary
            # estimation cache and cause warmup to hang or fail. Override it
            # for estimation, then restore it in configure_kv_cache_capacity().
            self._kv_cache_config.pool_ratio = None
            self._kv_cache_config.avg_seq_len = self._max_seq_len
            if self._is_kv_cache_manager_v2:
                free_mem, _ = torch.cuda.mem_get_info()
                max_gpu_total_bytes = int(
                    self._kv_cache_config.free_gpu_memory_fraction * free_mem)
                if (self._max_gpu_total_bytes_in is not None
                        and self._max_gpu_total_bytes_in > 0):
                    max_gpu_total_bytes = min(max_gpu_total_bytes,
                                              self._max_gpu_total_bytes_in)
                self._kv_cache_config.max_gpu_total_bytes = max_gpu_total_bytes
                self._kv_cache_config.max_tokens = max_tokens
            else:
                self._kv_cache_config.max_tokens = max_tokens
        return estimating_kv_cache

    def _configure_helix_kv_cache_capacity(self) -> None:
        """Set the helix KV quota without profiling (not CP-aware).

        Explicit quotas pass through; otherwise fraction sizing sets
        ``max_gpu_total_bytes`` (a rank-local byte cap the manager consumes
        as-is). Setting ``max_tokens`` here would overshoot the fraction:
        the manager inflates that knob by 1 / max_util_for_resume.
        """
        if (self._kv_cache_config.max_tokens is not None
                and self._kv_cache_config.max_tokens <= 0):
            raise ValueError(
                "Helix CP: kv_cache_config.max_tokens must be positive when "
                f"set, got {self._kv_cache_config.max_tokens}.")
        if (self._kv_cache_config.max_gpu_total_bytes or 0) > 0 or \
                (self._kv_cache_config.max_tokens or 0) > 0:
            logger.info("Helix CP: skipping KV cache capacity profiling; using "
                        "the explicitly configured quota.")
            return
        fraction = self._kv_cache_config.free_gpu_memory_fraction
        free_mem, _total = torch.cuda.mem_get_info()
        budget_bytes = int(free_mem * fraction)
        if budget_bytes <= 0:
            raise ValueError(
                "Helix CP: fraction-based KV sizing found no usable free "
                "memory; set kv_cache_config.max_tokens or "
                "max_gpu_total_bytes.")
        logger.warning(
            "Helix CP: capacity profiling is unsupported; sizing the KV "
            f"cache as fraction {fraction} of free memory -> "
            f"max_gpu_total_bytes={budget_bytes} (rank-local byte cap; the "
            "manager min-syncs across ranks and converts to global tokens). "
            "Set kv_cache_config.max_tokens or max_gpu_total_bytes to "
            "override.")
        self._kv_cache_config.max_gpu_total_bytes = budget_bytes

    def configure_kv_cache_capacity(self,
                                    py_executor: PyExecutor = None) -> None:
        """Perform KV cache capacity estimation.
        NOTE: for VSWA case, we calculate and set kv cache memory instead of using max_tokens in kv_cache_config.

        This updates `kv_cache_config`.
        """
        mapping = self._mapping

        # TODO: support CP by generating dummy requests for it.
        if mapping.cp_config.get('cp_type') == CpType.HELIX:
            if not self._is_kv_cache_manager_v2:
                # The helix sizing below emits V2 ledger (global) quotas;
                # V1 reads max_tokens as rank-local. Reject explicitly.
                raise NotImplementedError(
                    "TRTLLM_SKIP_KV_CACHE_ESTIMATION with helix CP requires "
                    "the V2 KV cache manager "
                    "(kv_cache_config.use_kv_cache_manager_v2=True).")
            self._configure_helix_kv_cache_capacity()
            return
        assert 'cp_type' not in mapping.cp_config

        fraction = self._kv_cache_config.free_gpu_memory_fraction

        torch.cuda.empty_cache()
        torch.cuda.reset_peak_memory_stats()
        end, total_gpu_memory = torch.cuda.mem_get_info()
        total_used_bytes = total_gpu_memory - end
        model_bytes = torch.cuda.memory_stats()["allocated_bytes.all.current"]
        logger.info(
            f"Memory used after loading model weights (inside torch) in memory usage profiling: {model_bytes / (GB):.2f} GiB"
        )
        logger.info(
            f"Memory used after loading model weights (outside torch) in memory usage profiling: {((total_used_bytes - model_bytes) if total_used_bytes > model_bytes else 0) / (GB):.2f} GiB"
        )

        profiled_output_bytes = 0

        if py_executor is not None and not self._skip_est:
            # Run the MM encoder at its independent token budget, then keep the
            # resulting request-owned embeddings resident while the text-only
            # LLM dummy fills max_num_tokens.
            encoder_profile_output = self._encode_dummy_inputs()
            if encoder_profile_output is not None:
                profiled_output_bytes = (encoder_profile_output.numel() *
                                         encoder_profile_output.element_size())
            py_executor.set_gather_responses(True)
            origin_iter_stats = py_executor.enable_iter_perf_stats
            py_executor.enable_iter_perf_stats = False
            req_ids = []
            if py_executor.dist.mapping.rank == 0:
                req_ids = py_executor.enqueue_requests(self._dummy_reqs)
            req_ids = py_executor.dist.broadcast(req_ids, root=0)
            py_executor.is_warmup = True
            py_executor.start_worker()
            try:
                responses = py_executor.await_responses(req_ids)
                for response_or_list in responses:
                    response_list = [response_or_list] if isinstance(
                        response_or_list,
                        ExecutorResponse) else response_or_list
                    for response in response_list:
                        if response.has_error():
                            raise RuntimeError(response.error_msg)

                torch_peak_memory = torch.cuda.memory_stats(
                )["allocated_bytes.all.peak"]

                # Release before measuring current usage so the retained
                # embeddings count toward the peak but not the steady state.
                encoder_profile_output = None

                # Clear the caching allocator before measuring the current memory usage
                torch.cuda.empty_cache()
                end, total_gpu_memory = torch.cuda.mem_get_info()
                torch_used_bytes = torch.cuda.memory_stats(
                )["allocated_bytes.all.current"]
            finally:
                # Redundant on the success path, but a failed dummy run would
                # otherwise keep the profiling embeddings alive through
                # teardown -- exactly when memory is already scarce.
                encoder_profile_output = None
                # get kv cache stats for both model and draft model
                kv_stats = py_executor.resource_manager.resource_managers.get(
                    ResourceManagerType.KV_CACHE_MANAGER).get_kv_cache_stats()
                # Get draft KV cache stats if present (either from two-model mode or one-model
                # mode with separate draft KV cache)
                draft_kv_cache_manager = py_executor.resource_manager.resource_managers.get(
                    ResourceManagerType.DRAFT_KV_CACHE_MANAGER)
                kv_stats_draft = draft_kv_cache_manager.get_kv_cache_stats(
                ) if draft_kv_cache_manager is not None else None

                # get total allocated bytes
                allocated_bytes = kv_stats.allocated_bytes + (
                    kv_stats_draft.allocated_bytes
                    if kv_stats_draft is not None else 0)
                py_executor.is_warmup = False
                py_executor.shutdown()
                py_executor.enable_iter_perf_stats = origin_iter_stats
                py_executor.set_gather_responses(False)

            total_used_bytes = total_gpu_memory - end
            activation_bytes = torch_peak_memory - model_bytes
            extra_cost = max(total_used_bytes - torch_used_bytes, 0)
            peak_memory = torch_peak_memory + extra_cost
            logger.info(
                f"Memory dynamically allocated during inference (inside torch) in memory usage profiling: {activation_bytes / (GB):.2f} GiB"
            )
            logger.info(
                f"Memory used outside torch (e.g., NCCL and CUDA graphs) in memory usage profiling: {extra_cost / (GB):.2f} GiB"
            )

        else:
            peak_memory = total_used_bytes
            allocated_bytes = 0
            activation_bytes = 0

        multimodal_encoder_memory_reserve = (
            self._get_multimodal_encoder_memory_reserve(
                profiled_output_bytes=profiled_output_bytes))
        peak_memory += multimodal_encoder_memory_reserve
        if multimodal_encoder_memory_reserve > 0:
            mem_gb = multimodal_encoder_memory_reserve / GB
            logger.info(
                f"Reserving {mem_gb:.2f} GiB for multimodal encoder memory "
                "not materialized by the profiling run.")

        # calculate max memory from peak memory and free gpu memory fraction
        kv_cache_max_memory = self._cal_max_memory(peak_memory,
                                                   total_gpu_memory, fraction,
                                                   allocated_bytes)

        # Estimation uses inferred pool sizing; the final manager uses the
        # user-provided configuration.
        self._kv_cache_config.pool_ratio = self._pool_ratio_in
        self._kv_cache_config.avg_seq_len = self._avg_seq_len_in

        # Reserve headroom for attention workspace the selected backend declares and the profiling forward
        # under-measures. KV-cache reuse can push summed attended KV past the profiled floor. Chunked prefill
        # usually bounds each attention launch by its chunk buffer, except for implementations such as the
        # NVFP4 DSA context gather that consume the complete attended prefix. The backend declares that
        # distinction through runtime_workspace_is_chunked_prefill_bounded. When a reserve applies,
        # reserve w * L_cap bytes -- covering the worst-case summed attended KV the scheduler admits
        # (get_mla_context_workspace_kv_len_cap) -- but clamp it to the per-token split budget * w / (k + w)
        # so a memory-constrained node shares the budget at a common token count instead of starving the
        # pool. Equivalently the pool keeps max((budget - w*L_cap)/k, budget/(k+w)) tokens. The reserve
        # covers exactly reserve/w tokens of summed attended KV; that count is carried to the KV manager as
        # the scheduler's admission cap so it never re-derives the cap from pool layout (which V2
        # overstates). No cap or w == 0 -> no-op.
        # A completed chunk-aware profiling run already prices the cached-KV
        # staging buffers in the measured peak. Do not subtract an additional
        # workspace reserve or install an attended-KV admission cap for it.
        profiled_mla_chunks = (py_executor is not None and not self._skip_est
                               and self._mla_chunked_profile_length is not None)
        w_bytes_per_token = (0 if profiled_mla_chunks else
                             get_attention_workspace_bytes_per_token(
                                 self._model_engine.model.model_config,
                                 self._mapping))
        workspace_is_chunked_prefill_bounded = True
        if w_bytes_per_token > 0:
            workspace_is_chunked_prefill_bounded = (
                get_attention_workspace_is_chunked_prefill_bounded(
                    self._model_engine.model.model_config))
        kv_len_cap = get_mla_context_workspace_kv_len_cap(
            self._kv_cache_config,
            self._max_batch_size,
            self._max_num_tokens,
            self._max_seq_len,
            self._llm_args.enable_chunked_prefill,
            workspace_is_chunked_prefill_bounded,
            chunked_workspace_profiled=profiled_mla_chunks,
            # Chunk-aware profiling is SM100+ only. Preserve Hopper's existing
            # backend-declared policy rather than adding a new reserve/cap to
            # its dense full-gather path just because it cannot profile chunks.
            require_chunked_workspace_profile=(
                w_bytes_per_token > 0 and self._llm_args.enable_chunked_prefill
                and get_sm_version() >= 100))
        if w_bytes_per_token > 0 and kv_len_cap:
            budget_before = kv_cache_max_memory
            workspace_reserve, self._fp8_ctx_mla_kv_len_cap = (
                get_mla_context_workspace_reserve(
                    budget_before,
                    self._get_kv_size_per_token().slope, w_bytes_per_token,
                    kv_len_cap))
            if workspace_reserve > 0:
                kv_cache_max_memory = int(budget_before - workspace_reserve)
                logger.info(
                    f"Reserving {workspace_reserve / (GB):.2f} GiB for the context-MLA attention "
                    f"workspace (w={w_bytes_per_token} B/token, admitting up to "
                    f"{self._fp8_ctx_mla_kv_len_cap} tokens of summed attended KV): KV cache budget "
                    f"{budget_before / (GB):.2f} -> {kv_cache_max_memory / (GB):.2f} GiB."
                )

        # NOTE:
        # For KVCacheManager, KvCacheCreator currently controls capacity using two parameters in KVCacheConfig:
        #   • max_tokens
        #   • max_gpu_total_bytes
        # For KVCacheManagerV2, KvCacheCreator controls capacity using max_gpu_total_bytes only.
        # This leaves max_tokens as a user-defined constraint.

        # ---------------------------handle max_tokens---------------------------------
        if self._is_kv_cache_manager_v2:
            # KVCacheManagerV2 doesn't rely on max_tokens to control capacity, so restore user provided value
            self._kv_cache_config.max_tokens = self._max_kv_tokens_in
        else:
            # handle user provided max_tokens
            if self._max_kv_tokens_in is not None:
                # raise error if it is VSWA case
                is_vswa = uses_vswa_kv_cache_layout(
                    self._kv_cache_config.max_attention_window)

                # raise error if it is VSWA case
                if is_vswa:
                    logger.warning(
                        "max_tokens should not be set for VSWA case as it is ambiguous concept for VSWA."
                    )
                # calculate max memory from max_tokens
                kv_size_per_token = self._get_kv_size_per_token()
                kv_cache_max_memory_from_max_tokens = (
                    kv_size_per_token.bytes_for_tokens(self._max_kv_tokens_in))
                kv_cache_max_memory = min(kv_cache_max_memory,
                                          kv_cache_max_memory_from_max_tokens)
                logger.info(
                    f"max_tokens={self._max_kv_tokens_in} is provided. It limits max memory to {kv_cache_max_memory_from_max_tokens / (GB):.2f} GiB. "
                    f"New max_memory is set to {kv_cache_max_memory / (GB):.2f} GiB"
                )
            # For KvCacheManager, its logic still relies on max_tokens to control capacity
            self._kv_cache_config.max_tokens = (self._get_kv_size_per_token(
            ).tokens_for_budget(kv_cache_max_memory))
        # ---------------------------handle max_tokens---------------------------------

        # ---------------------------handle max_gpu_total_bytes---------------------------------
        # if user provided max_gpu_total_bytes, set max memory from max_gpu_total_bytes
        if (self._max_gpu_total_bytes_in is not None
                and self._max_gpu_total_bytes_in > 0):
            kv_cache_max_memory = min(kv_cache_max_memory,
                                      self._max_gpu_total_bytes_in)
            logger.info(
                f"max_gpu_total_bytes={self._max_gpu_total_bytes_in / (GB):.2f} GiB is provided. New max memory is {kv_cache_max_memory / (GB):.2f} GiB"
            )

        logger.info(
            f"Estimated max memory in KV cache : {kv_cache_max_memory / (GB):.2f} GiB"
        )
        # set max_gpu_total_bytes
        self._kv_cache_config.max_gpu_total_bytes = kv_cache_max_memory
        if isinstance(self._profiling_stage_data, dict):
            self._profiling_stage_data["activation_bytes"] = activation_bytes
        # ---------------------------handle max_gpu_total_bytes---------------------------------

    def _create_kv_cache_manager(
        self,
        model_engine: PyTorchModelEngine,
        estimating_kv_cache: bool = False,
        kv_cache_config_override: Optional[KvCacheConfig] = None,
        cold_page_codec_provider: Optional[object] = None,
    ) -> KVCacheManager:
        mapping = self._mapping
        assert model_engine.model.model_config.is_generation, "Only construct KV cache for generation models."
        kv_cache_config = (kv_cache_config_override if kv_cache_config_override
                           is not None else self._kv_cache_config)
        kv_cache_manager_cls = self._get_model_kv_cache_manager_cls(
            model_engine, kv_cache_config)

        # When using separate draft KV cache in one-model speculative decoding,
        # use layer_mask to include only target layers. The draft layers should
        # only be in the separate draft KV cache manager.
        # We still pass spec_config so that num_extra_kv_tokens is calculated.
        spec_dec_layer_mask = None
        if self._should_create_separate_draft_kv_cache():
            num_target_layers = model_engine.model.model_config.pretrained_config.num_hidden_layers
            spec_dec_layer_mask = [True] * num_target_layers

        estimating_kv_cache = estimating_kv_cache and not self._skip_est
        kv_cache_manager = _create_kv_cache_manager(
            model_engine=model_engine,
            kv_cache_manager_cls=kv_cache_manager_cls,
            mapping=mapping,
            kv_cache_config=kv_cache_config,
            tokens_per_block=self._tokens_per_block,
            max_seq_len=self._max_seq_len,
            max_batch_size=self._max_batch_size,
            spec_config=self._speculative_config,
            sparse_attention_config=self._sparse_attention_config,
            max_num_tokens=self._max_num_tokens,
            max_beam_width=self._max_beam_width,
            kv_connector_manager=self._kv_connector_manager,
            estimating_kv_cache=estimating_kv_cache,
            enable_kv_cache_stats=self._enable_kv_cache_stats()
            and not estimating_kv_cache,
            execution_stream=self._execution_stream,
            layer_mask=spec_dec_layer_mask,
            is_disagg=self._is_disagg,
            disable_overlap_scheduler=self._disable_overlap_scheduler,
            kv_events_config=None
            if estimating_kv_cache or model_engine.is_draft_model else
            self._llm_args.kv_cache_config.kv_events_config,
            cold_page_codec_provider=cold_page_codec_provider,
            joint_kv_cache_reuse=self._joint_kv_cache_reuse,
        )

        if not self._skip_est:
            # KVCacheManager (Non-draft) modifies the max_seq_len field, update it to self
            if model_engine.kv_cache_manager_key == ResourceManagerType.KV_CACHE_MANAGER:
                # When SWA is enabled, max_seq_len is updated inside kv_cache_manager.
                if kv_cache_manager is not None:
                    if kv_cache_manager.max_seq_len < self._max_seq_len:
                        self._dummy_reqs = self._create_dummy_context_requests(
                            max(
                                1, self._net_max_seq_len - 1 -
                                (self._max_seq_len -
                                 kv_cache_manager.max_seq_len)))
                    self._max_seq_len = kv_cache_manager.max_seq_len

                # When SWA is enabled, max_seq_len is updated inside kv_cache_manager.
                if kv_cache_manager is not None:
                    if kv_cache_manager.max_seq_len < self._max_seq_len:
                        self._dummy_reqs = self._create_dummy_context_requests(
                            max(1, kv_cache_manager.max_seq_len - 1))
                    self._max_seq_len = kv_cache_manager.max_seq_len
        else:
            if kv_cache_manager is not None:
                self._max_seq_len = kv_cache_manager.max_seq_len

        return kv_cache_manager

    def _should_create_separate_draft_kv_cache(self) -> bool:
        """
        Check if we need a separate draft KV cache manager for one-model mode.
        Returns True if the speculative config has use_separate_draft_kv_cache=True.

        Note: For MTP, _draft_config may be None since MTP layers are embedded
        in the target model and don't produce a separate ModelConfig. We fall
        back to the target model's config via _get_effective_draft_config().
        """
        if self._speculative_config is None:
            # No drafter at all, so there is nothing to give a manager to.
            return False
        # Narrower than is_external_drafter(): PARD and DRAFT_TARGET_ONE_MODEL
        # never reach the arena this carve-out exists for.
        spec_dec_mode = self._speculative_config.spec_dec_mode
        is_standalone_drafter = (spec_dec_mode.is_dflash()
                                 or spec_dec_mode.is_dspark())
        if self._mapping.enable_attention_dp and not is_standalone_drafter:
            # This bail suits MTP, whose draft layers are target-shaped and
            # appendable to the target pool. A standalone drafter has nothing to
            # append, so it would be stranded on the private arena instead.
            logger.info(
                "Attention DP is enabled, separate draft KV cache is not supported."
            )
            return False

        sparse_cfg = self._sparse_attention_config
        if (sparse_cfg is not None
                and getattr(sparse_cfg, "algorithm", None) == "deepseek_v4"
                and self._mapping.pp_size > 1):
            logger.info(
                "DeepSeek-V4 separate draft KV cache is only supported for PP=1; "
                "folding draft layers into the unified manager for pp_size=%d.",
                self._mapping.pp_size)
            return False
        return should_use_separate_draft_kv_cache(self._speculative_config)

    def _joint_reuse_supported(self) -> bool:
        """Span half of the pairing decision; ``build_managers`` ANDs it with
        ``has_separate_one_model_draft``. ``False`` = leave the run unpaired.
        """
        if not self._is_kv_cache_manager_v2:
            return False
        lookahead = draft_prompt_lookahead(self._speculative_config)
        if lookahead is None:
            return False
        if lookahead > 0 and not getattr(self._kv_cache_manager_cls,
                                         "_supports_reuse_match_backoff",
                                         False):
            # The opt-out is about backing the match off by `lookahead` tokens,
            # which a specialized commit/history protocol (recurrent snapshots,
            # DSA) cannot express. A zero span asks for no backoff at all, so
            # every backoff-sized path stays a no-op and the pairing is safe.
            return False
        return True

    def _get_effective_draft_config(self) -> ModelConfig:
        """
        Return the ModelConfig to use for draft KV cache creation.

        For Eagle3 and draft-target one-model modes, a dedicated draft config
        is provided at construction time.  For MTP one-model mode, no separate
        draft config exists because the MTP layers share the same architecture
        as the target model.  In that case we fall back to the target model's
        config so that the draft KV cache is created with the correct layout.
        """
        if self._draft_config is not None:
            return self._draft_config
        # MTP: MTP layers reuse the target model architecture, so the target
        # model's config describes the correct KV cache layout for the draft
        # layers as well.
        return self._model_engine.model.model_config

    def _get_draft_kv_model_config(self) -> ModelConfig:
        """The draft ModelConfig describing the KV pool as it is ALLOCATED.

        The args-level ``kv_cache_config.dtype`` sync stamps the TARGET's fp8 KV
        algo onto every loaded model, including a standalone drafter. The drafter
        stores and reads its pool in its weights dtype (DFlash validates a bf16 pool
        and otherwise falls back to the max_seq_len-dense private arena, which OOMs at
        long context), so the pool dtype must follow the drafter.

        Every consumer of draft KV bytes must go through here. If the budget split
        and the allocation read different dtypes, the split charges fp8 bytes for a
        bf16 pool and the draft manager gets HALF the target's tokens. The capacity
        scheduler admits on the target pool alone, so past ~50% target utilization it
        raises "Draft KV cache context resize failed", fatal to every rank.
        """
        effective_draft_config = self._get_effective_draft_config()
        # Narrower than is_external_drafter(), matching
        # _should_create_separate_draft_kv_cache. PARD and DRAFT_TARGET_ONE_MODEL
        # reach here too and can carry a genuine fp8 KV algo of their own, which
        # dtype="auto" keeps; dropping it would allocate bf16 under attention
        # modules that still read and write fp8.
        spec_dec_mode = self._speculative_config.spec_dec_mode
        if not (spec_dec_mode.is_dflash() or spec_dec_mode.is_dspark()):
            return effective_draft_config
        quant_config = getattr(effective_draft_config, "quant_config", None)
        if quant_config is None or not quant_config.quant_mode.has_fp8_kv_cache(
        ):
            return effective_draft_config
        logger.info(
            "External drafter KV pool keeps the drafter dtype; dropping "
            "the fp8 KV quant algo inherited from the target.")
        neutral_quant = copy.copy(quant_config)
        neutral_quant.kv_cache_quant_algo = None
        # QuantConfig.quant_mode and .layer_quant_mode are both cached_property
        # and the copy carries the already-computed caches, so BOTH must be
        # dropped for the mutation to take: _create_kv_cache_manager reads
        # quant_mode off this copy, and layer_quant_mode is the pair's other
        # half, stale in the same way.
        neutral_quant.__dict__.pop("quant_mode", None)
        neutral_quant.__dict__.pop("layer_quant_mode", None)
        # No _frozen dance: ModelConfig.__setattr__ exempts quant_config by
        # name, and restoring _frozen to True would freeze a copy whose source
        # may not have been frozen.
        effective_draft_config = copy.copy(effective_draft_config)
        effective_draft_config.quant_config = neutral_quant
        return effective_draft_config

    def _get_num_draft_layers(self) -> int:
        """Return the actual number of draft KV cache layers.

        This must stay in sync with the num_layers passed to the draft KV
        cache manager constructor in _create_one_model_draft_kv_cache_manager.
        """
        if self._speculative_config.spec_dec_mode.is_external_drafter():
            return self._draft_config.pretrained_config.num_hidden_layers
        return get_num_spec_layers(self._speculative_config)

    def _get_draft_max_attention_window(
        self,
        max_seq_len: int,
        kv_cache_config: KvCacheConfig,
    ) -> Optional[List[int]]:
        """Derive the draft manager's per-layer attention windows."""
        effective_draft_config = self._get_effective_draft_config()
        return _derive_draft_max_attention_window(
            kv_cache_config,
            effective_draft_config.pretrained_config,
            max_seq_len,
            self._get_num_draft_layers(),
        )

    def _get_one_model_draft_kv_cache_config(
        self,
        kv_cache_config: KvCacheConfig,
        max_seq_len: int,
        *,
        estimating_kv_cache: bool = False,
    ) -> KvCacheConfig:
        """Return a clone with the draft manager's attention-window layout."""
        # Estimation uses a small max_tokens-sized temporary draft cache before
        # the measured GPU budget is available to split. Applying VSWA there
        # would size every window pool from the unsplit free-memory budget.
        max_attention_window = (None if estimating_kv_cache else
                                self._get_draft_max_attention_window(
                                    max_seq_len, kv_cache_config))
        return kv_cache_config.model_copy(
            update={"max_attention_window": max_attention_window})

    def _create_one_model_draft_kv_cache_manager(
        self,
        max_seq_len: int,
        estimating_kv_cache: bool = False,
        kv_cache_config_override: Optional[KvCacheConfig] = None,
        cold_page_codec_provider: Optional[object] = None,
    ) -> Optional[KVCacheManager]:
        """
        Create a KV cache manager for draft model layers in one-model mode
        when target and draft have different KV cache layouts.
        """
        num_draft_layers = self._get_num_draft_layers()
        spec_dec_layer_mask = self._get_one_model_draft_layer_mask()

        # Get the effective draft config (explicit draft_config if available,
        # otherwise fall back to target model config for MTP), with the
        # target's inherited fp8 KV algo dropped for a standalone drafter. The
        # budget split in _get_kv_size_per_token resolves it through the SAME
        # helper, so the bytes/token it charges match the pool allocated here.
        effective_draft_config = self._get_draft_kv_model_config()

        kv_cache_config = (kv_cache_config_override if kv_cache_config_override
                           is not None else self._kv_cache_config)
        draft_kv_config = self._get_one_model_draft_kv_cache_config(
            kv_cache_config,
            max_seq_len,
            estimating_kv_cache=estimating_kv_cache)
        draft_kv_config.max_attention_window = (
            _expand_attention_window_pattern_to_global_layers(
                draft_kv_config.max_attention_window,
                spec_dec_layer_mask,
            ))
        if (not uses_vswa_kv_cache_layout(draft_kv_config.max_attention_window)
                and draft_kv_config.pool_ratio is not None
                and len(draft_kv_config.pool_ratio) != 1):
            # pool_ratio describes one manager's layer-group layout. The
            # target hybrid manager may have separate recurrent-state and
            # attention layer groups, while a non-VSWA draft manager has one
            # attention layer group. Reusing the target's ratios fails its arity
            # check.
            logger.info(
                "Normalizing the separate one-model draft KV cache pool_ratio "
                f"from {draft_kv_config.pool_ratio} to [1.0] for its single "
                "layer group.")
            draft_kv_config.pool_ratio = [1.0]
        if uses_vswa_kv_cache_layout(draft_kv_config.max_attention_window):
            logger.info(
                f"Derived draft KV cache max_attention_window for separate "
                f"draft manager: {draft_kv_config.max_attention_window}")
        # Get the appropriate KV cache manager class for the draft model
        draft_kv_cache_manager_cls = get_kv_cache_manager_cls(
            effective_draft_config,
            draft_kv_config,
            is_disagg=self._is_disagg,
            cache_transceiver_config=self._cache_transceiver_config)
        draft_kv_cache_manager_cls = self._validate_or_fallback_kv_cache_manager_v2(
            draft_kv_cache_manager_cls, effective_draft_config, draft_kv_config)

        estimating_kv_cache = estimating_kv_cache and not self._skip_est
        # For MTP with models using sparse attention (e.g., DeepSeek V3 with DSA),
        # the draft layers share the same architecture as the target model and need
        # the sparse_attention_config. Get it from effective_draft_config which
        # falls back to the target model's config for MTP mode.
        sparse_attn_config = effective_draft_config.sparse_attention_config
        # Under helix the standalone drafter is built against the repurposed
        # mapping (CP ranks become TP ranks) and every rank keeps its full
        # drafter KV, so its paged manager needs the CP-free mapping: the
        # round-robin ledger applies to the TARGET KV alone, and
        # KVCacheManagerV2 rejects is_draft x helix outright.
        draft_mapping = self._mapping
        if draft_mapping.has_cp_helix():
            draft_mapping = draft_mapping.repurpose_helix_cp_to_tp()
        return _create_kv_cache_manager(
            model_engine=None,
            max_cuda_graph_batch_size=self._model_engine.
            _max_cuda_graph_batch_size,
            kv_cache_manager_cls=draft_kv_cache_manager_cls,
            mapping=draft_mapping,
            kv_cache_config=draft_kv_config,
            tokens_per_block=self._tokens_per_block,
            max_seq_len=max_seq_len,
            max_batch_size=self._max_batch_size,
            spec_config=self._speculative_config,
            sparse_attention_config=sparse_attn_config,
            max_num_tokens=self._max_num_tokens,
            max_beam_width=self._max_beam_width,
            kv_connector_manager=self._kv_connector_manager,
            estimating_kv_cache=estimating_kv_cache,
            enable_kv_cache_stats=self._enable_kv_cache_stats()
            and not estimating_kv_cache,
            execution_stream=self._execution_stream,
            # One-model draft specific overrides
            model_config=effective_draft_config,
            dtype=effective_draft_config.pretrained_config.torch_dtype,
            is_draft=True,
            layer_mask=spec_dec_layer_mask,
            num_layers=num_draft_layers,
            is_disagg=self._is_disagg,
            disable_overlap_scheduler=self._disable_overlap_scheduler,
            cold_page_codec_provider=cold_page_codec_provider,
            joint_kv_cache_reuse=self._joint_kv_cache_reuse,
        )

    def _get_target_and_draft_cache_costs(
        self,
        kv_cache_config: Optional[KvCacheConfig] = None,
    ) -> Optional[tuple[CacheCost, CacheCost]]:
        """Per-manager KV cache costs for target and draft layers."""
        target_kv_cache_config = (kv_cache_config if kv_cache_config is not None
                                  else self._kv_cache_config)
        use_separate_draft_kv_cache = (
            self._should_create_separate_draft_kv_cache())
        target_kv = self._per_manager_cache_cost(
            self._kv_cache_manager_cls,
            self._model_engine.model.model_config,
            target_kv_cache_config,
            use_separate_draft_kv_cache=use_separate_draft_kv_cache)
        # Estimate the draft component directly so its independently modelled
        # affine intercept is preserved exactly.
        draft_kv = self._get_draft_cache_cost(
            target_kv_cache_config,
            use_separate_draft_kv_cache=use_separate_draft_kv_cache,
        )
        if draft_kv is None:
            return None
        costs = (target_kv, draft_kv)
        if any(cost.slope < 0 or cost.intercept < 0 or (
                cost.slope == 0 and cost.intercept == 0) for cost in costs):
            return None
        return target_kv, draft_kv

    def _compute_draft_budget_shares(
        self,
        total_budget: int,
        target_kv: CacheCost,
        draft_kv: CacheCost,
    ) -> Optional[tuple[int, int]]:
        """Split *total_budget* into (target_budget, draft_budget) byte shares."""
        intercept_total = target_kv.intercept + draft_kv.intercept
        slope_budget = total_budget - intercept_total
        slope_total = target_kv.slope + draft_kv.slope
        if slope_budget < 0:
            logger.warning(
                f"KV cache budget {total_budget} is smaller than the fixed "
                f"cache cost {intercept_total}; cannot split between "
                f"target and draft.")
            return None
        if slope_budget == 0 and slope_total > 0:
            logger.warning(
                f"KV cache budget {total_budget} leaves no capacity beyond "
                f"the fixed cache cost {intercept_total}; cannot split "
                f"between target and draft with a per-token cache cost.")
            return None
        draft_slope_share = (slope_budget * draft_kv.slope //
                             slope_total if slope_total > 0 else 0)
        draft_budget = draft_kv.intercept + draft_slope_share
        target_budget = total_budget - draft_budget
        return target_budget, draft_budget

    def _split_kv_cache_budget_for_draft(
        self,
        budget_attr: str,
        target_kv_cache_config: Optional[KvCacheConfig] = None,
        draft_kv_cache_config: Optional[KvCacheConfig] = None,
    ) -> tuple[KvCacheConfig, Optional[KvCacheConfig]]:
        """Split a byte budget (attribute on ``KvCacheConfig``) between target
        and draft KV caches.

        Splits the value of ``target_kv_cache_config.<budget_attr>`` using the
        affine target/draft cache costs, then returns cloned target and draft
        configs containing their respective shares.

        The input target config and the creator's base config are not mutated.
        When the split is not applicable (the budget is not set, or the
        per-manager cache costs are unavailable), the input configs are returned
        unchanged.

        The affine fixed (intercept) cost models GPU-resident state (e.g. mamba
        SSM state). It is only charged against ``max_gpu_total_bytes``; for any
        other budget (e.g. ``host_cache_size``, which is host offload memory the
        GPU-resident state never occupies) the intercept is dropped so the split
        stays proportional to the per-token cost.

        When the split is *infeasible* (the combined fixed cost exhausts the
        budget while either manager has a per-token cost, or exceeds it)
        the shortfall is fatal: both managers need their fixed state resident in
        GPU memory, so the run would OOM. It raises ``ValueError`` rather than
        silently producing an unusable config. A defensive degrade-to-zero path
        for non-GPU budgets remains so the draft never silently inherits the full
        budget and double-allocates it.
        """
        target_kv_cache_config = (target_kv_cache_config
                                  if target_kv_cache_config is not None else
                                  self._kv_cache_config)
        total_budget = getattr(target_kv_cache_config, budget_attr) or 0
        if total_budget <= 0:
            return target_kv_cache_config, draft_kv_cache_config

        cache_costs = self._get_target_and_draft_cache_costs(
            target_kv_cache_config)
        if cache_costs is None:
            return target_kv_cache_config, draft_kv_cache_config
        target_kv, draft_kv = cache_costs

        # The fixed (intercept) cost models GPU-resident state such as mamba SSM
        # state; it does not consume host offload memory. When splitting a
        # non-GPU budget (e.g. host_cache_size), drop the intercept so the split
        # stays proportional to the per-token (slope) cost instead of being
        # spuriously starved by a GPU-only fixed cost.
        if budget_attr != "max_gpu_total_bytes":
            target_kv = CacheCost(slope=target_kv.slope)
            draft_kv = CacheCost(slope=draft_kv.slope)

        shares = self._compute_draft_budget_shares(total_budget, target_kv,
                                                   draft_kv)
        if shares is None:
            # The split cannot provide each manager with usable GPU capacity.
            intercept_total = target_kv.intercept + draft_kv.intercept
            if budget_attr == "max_gpu_total_bytes":
                # A GPU budget that cannot even fit the combined fixed cost is
                # fatal: both managers need their fixed state resident in GPU
                # memory, so the run would OOM. Fail fast with actionable
                # guidance rather than producing an unusable zero-budget draft.
                raise ValueError(
                    f"KV cache GPU budget ({total_budget / GB:.2f} GiB) is "
                    f"insufficient after the combined fixed cost "
                    f"({intercept_total / GB:.2f} GiB, e.g. SWA or mamba state) "
                    f"for target+draft. Increase free_gpu_memory_fraction or "
                    f"max_gpu_total_bytes, or reduce max_batch_size (the fixed "
                    f"cost scales with batch size).")
            # Defensive: non-GPU budgets zero out the intercept above, so with a
            # positive budget this branch is currently unreachable for them. It
            # remains as a safety net guaranteeing that, should a non-GPU budget
            # ever carry a fixed cost it cannot fit, we degrade gracefully rather
            # than letting both managers inherit the full budget and
            # double-allocate it: keep the full budget on the target and zero the
            # draft's share for this attribute.
            logger.warning(
                f"Cannot split KV cache {budget_attr} between target and draft; "
                f"assigning the draft a zero {budget_attr} budget to avoid "
                f"double-allocating the full budget.")
            if draft_kv_cache_config is None:
                draft_kv_cache_config = target_kv_cache_config.model_copy()
            else:
                draft_kv_cache_config = draft_kv_cache_config.model_copy()
            setattr(draft_kv_cache_config, budget_attr, 0)
            return target_kv_cache_config, draft_kv_cache_config
        target_budget, draft_budget = shares

        logger.info(
            f"Splitting KV cache {budget_attr}: total={total_budget / GB:.2f} GiB, "
            f"target={target_budget / GB:.2f} GiB ({target_kv}), "
            f"draft={draft_budget / GB:.2f} GiB ({draft_kv})")

        split_target_kv_cache_config = target_kv_cache_config.model_copy()
        setattr(split_target_kv_cache_config, budget_attr, target_budget)
        if draft_kv_cache_config is None:
            split_draft_kv_cache_config = target_kv_cache_config.model_copy()
        else:
            split_draft_kv_cache_config = draft_kv_cache_config.model_copy()
        setattr(split_draft_kv_cache_config, budget_attr, draft_budget)
        return split_target_kv_cache_config, split_draft_kv_cache_config

    def _is_encoder_decoder(self) -> bool:
        return self._model_engine.model.model_config.is_encoder_decoder

    @staticmethod
    def _get_config_int_attr(config, names: tuple[str, ...]) -> Optional[int]:
        for name in names:
            value = getattr(config, name, None)
            if isinstance(value, int):
                return value
        return None

    def _get_cross_kv_cache_layout(
        self,
        fallback_max_seq_len: Optional[int] = None
    ) -> tuple[int, int, int, int]:
        """Return decoder-layer count and encoder KV geometry for cross cache."""
        config = self._model_engine.model.model_config.pretrained_config

        num_layers = self._get_config_int_attr(
            config,
            ("num_decoder_layers", "decoder_layers", "num_hidden_layers",
             "num_layers"),
        )
        if num_layers is None:
            raise ValueError(
                "Unable to determine decoder layer count for cross KV cache.")

        encoder_num_heads = self._get_config_int_attr(
            config,
            ("encoder_num_heads", "encoder_attention_heads", "num_heads",
             "num_attention_heads"),
        )
        if encoder_num_heads is None:
            raise ValueError(
                "Unable to determine encoder attention head count for cross KV cache."
            )

        num_kv_heads = self._get_config_int_attr(
            config,
            ("encoder_num_kv_heads", "encoder_num_key_value_heads",
             "encoder_attention_heads", "encoder_num_heads",
             "num_key_value_heads", "num_heads", "num_attention_heads"),
        )
        if num_kv_heads is None:
            num_kv_heads = encoder_num_heads

        encoder_hidden_size = self._get_config_int_attr(
            config, ("encoder_hidden_size", "d_model", "hidden_size"))
        if encoder_hidden_size is None:
            raise ValueError(
                "Unable to determine encoder hidden size for cross KV cache.")

        head_dim = self._get_config_int_attr(
            config,
            ("encoder_head_size", "encoder_head_dim", "d_kv"),
        )
        if head_dim is None:
            head_dim = encoder_hidden_size // encoder_num_heads

        max_seq_len = fallback_max_seq_len or self._max_seq_len
        max_input_len = getattr(self._llm_args, "max_input_len", None)
        if isinstance(max_input_len, int) and max_input_len > 0:
            max_seq_len = max_input_len
        encoder_limit = self._get_config_int_attr(
            config,
            ("max_encoder_input_len", "encoder_max_input_length",
             "max_encoder_position_embeddings",
             "encoder_max_position_embeddings", "max_position_embeddings",
             "n_positions"),
        )
        if encoder_limit is not None:
            max_seq_len = min(max_seq_len, encoder_limit)

        return num_layers, num_kv_heads, head_dim, max_seq_len

    def _get_cross_kv_size_per_token(self) -> int:
        """Estimate bytes/token for the encoder-decoder cross-attention pool."""
        from types import SimpleNamespace

        model_config = self._model_engine.model.model_config
        config = model_config.pretrained_config
        (num_layers, num_kv_heads, head_dim,
         _) = self._get_cross_kv_cache_layout()
        num_attention_heads = self._get_config_int_attr(
            config,
            ("encoder_num_heads", "encoder_attention_heads", "num_heads",
             "num_attention_heads"),
        )
        hidden_size = self._get_config_int_attr(
            config, ("encoder_hidden_size", "d_model", "hidden_size"))
        proxy_model_config = SimpleNamespace(
            pretrained_config=SimpleNamespace(
                num_key_value_heads=num_kv_heads,
                num_attention_heads=num_attention_heads,
                hidden_size=hidden_size,
                head_dim=head_dim,
            ),
            quant_config=model_config.quant_config,
        )
        return self._kv_cache_manager_cls.get_cache_size_per_token(
            proxy_model_config,
            self._mapping,
            tokens_per_block=self._tokens_per_block,
            num_layers=num_layers,
        )

    def _split_kv_cache_budget_for_cross(
        self,
        kv_cache_config: Optional[KvCacheConfig] = None,
    ) -> tuple[KvCacheConfig, KvCacheConfig]:
        """Split enc-dec KV cache budgets between self and cross pools.

        The cross manager must exist for every encoder-decoder runtime. During
        both estimation and final construction, split the same memory-derived
        budget sources used by the legacy TRT path: the free-memory fraction,
        any explicit ``max_gpu_total_bytes`` override, and any explicit host
        cache budget. ``max_tokens`` is a logical cap, not a memory split knob,
        so it is intentionally left unchanged. The creator's base config is not
        mutated.
        """
        base_kv_cache_config = (kv_cache_config if kv_cache_config is not None
                                else self._kv_cache_config)
        fraction = base_kv_cache_config.cross_kv_cache_fraction
        if fraction is None:
            raise ValueError("Encoder-decoder models require "
                             "cross_kv_cache_fraction to size the cross "
                             "KV cache pool.")

        self_kv_cache_config = base_kv_cache_config.model_copy()
        cross_kv_cache_config = base_kv_cache_config.model_copy()
        split_any_budget = False

        free_fraction = base_kv_cache_config.free_gpu_memory_fraction
        if free_fraction is not None:
            cross_fraction = free_fraction * fraction
            self_fraction = free_fraction - cross_fraction
            logger.info(
                "Splitting encoder-decoder free GPU memory fraction: "
                f"total={free_fraction:.3f}, self={self_fraction:.3f}, cross={cross_fraction:.3f}"
            )
            self_kv_cache_config.free_gpu_memory_fraction = self_fraction
            cross_kv_cache_config.free_gpu_memory_fraction = cross_fraction
            split_any_budget = True

        total_budget = base_kv_cache_config.max_gpu_total_bytes
        if total_budget is not None and total_budget > 0:
            cross_budget = int(total_budget * fraction)
            self_budget = total_budget - cross_budget
            logger.info(
                f"Splitting KV cache budget for encoder-decoder: "
                f"total={total_budget / GB:.2f} GiB, "
                f"self={self_budget / GB:.2f} GiB ({1 - fraction:.0%}), "
                f"cross={cross_budget / GB:.2f} GiB ({fraction:.0%})")
            self_kv_cache_config.max_gpu_total_bytes = self_budget
            cross_kv_cache_config.max_gpu_total_bytes = cross_budget
            split_any_budget = True

        host_cache_size = base_kv_cache_config.host_cache_size
        if host_cache_size is not None and host_cache_size > 0:
            cross_host_cache_size = int(host_cache_size * fraction)
            self_host_cache_size = host_cache_size - cross_host_cache_size
            logger.info(
                f"Splitting KV cache host budget for encoder-decoder: "
                f"total={host_cache_size / GB:.2f} GiB, "
                f"self={self_host_cache_size / GB:.2f} GiB ({1 - fraction:.0%}), "
                f"cross={cross_host_cache_size / GB:.2f} GiB ({fraction:.0%})")
            self_kv_cache_config.host_cache_size = self_host_cache_size
            cross_kv_cache_config.host_cache_size = cross_host_cache_size
            split_any_budget = True

        if not split_any_budget:
            raise ValueError("Unable to size the encoder-decoder cross KV "
                             "cache pool: neither free_gpu_memory_fraction nor "
                             "max_gpu_total_bytes nor host_cache_size is "
                             "available.")

        return self_kv_cache_config, cross_kv_cache_config

    def _create_cross_kv_cache_manager(
        self,
        cross_kv_cache_config: KvCacheConfig,
        estimating_kv_cache: bool = False,
        fallback_max_seq_len: Optional[int] = None,
    ) -> KVCacheManager:
        """Create a KV cache manager for the cross-attention pool.

        The cross pool stores encoder K/V projections that are written once
        during the first decoder context step and read on every subsequent
        decoder generation step. It uses ``CacheType.CROSS`` with decoder
        layer count but encoder-side KV geometry.

        The manager class mirrors the self pool (``KVCacheManager`` for V1,
        ``KVCacheManagerV2`` for V2) so that both pools share the same
        runtime ABI and scheduler integration. V1 is the default and the
        production target for encoder-decoder models.
        """
        (num_layers, num_kv_heads, head_dim,
         max_seq_len) = self._get_cross_kv_cache_layout(fallback_max_seq_len)
        estimating_kv_cache = estimating_kv_cache and not self._skip_est
        return _create_kv_cache_manager(
            model_engine=self._model_engine,
            kv_cache_manager_cls=self._kv_cache_manager_cls,
            mapping=self._mapping,
            kv_cache_config=cross_kv_cache_config,
            tokens_per_block=self._tokens_per_block,
            max_seq_len=max_seq_len,
            max_batch_size=self._max_batch_size,
            spec_config=None,
            sparse_attention_config=None,
            max_num_tokens=self._max_num_tokens,
            max_beam_width=1,
            kv_connector_manager=None,
            estimating_kv_cache=estimating_kv_cache,
            execution_stream=self._execution_stream,
            num_layers=num_layers,
            num_kv_heads=num_kv_heads,
            head_dim=head_dim,
            disable_overlap_scheduler=self._disable_overlap_scheduler,
            kv_cache_type=tensorrt_llm.bindings.internal.batch_manager.
            CacheType.CROSS,
        )

    def _needs_gpu_kv_cache_budget_split(
        self,
        max_seq_len: int,
        kv_cache_config: Optional[KvCacheConfig] = None,
    ) -> bool:
        """Whether max_gpu_total_bytes must be split per manager."""
        if self._is_kv_cache_manager_v2:
            return self._should_create_separate_draft_kv_cache()
        kv_cache_config = (kv_cache_config if kv_cache_config is not None else
                           self._kv_cache_config)
        if uses_vswa_kv_cache_layout(kv_cache_config.max_attention_window):
            return True
        if not self._should_create_separate_draft_kv_cache():
            return False
        draft_windows = self._get_draft_max_attention_window(
            max_seq_len, kv_cache_config)
        return uses_vswa_kv_cache_layout(draft_windows)

    @classmethod
    def _drop_explicit_offload_tier_budgets(
            cls, kv_cache_config: Optional[KvCacheConfig]
    ) -> Optional[KvCacheConfig]:
        """Return a copy of the config with explicit offload budgets unset.

        Host sizing then falls to the V2 auto host tier policy, which matches
        the tier to the manager's own device quota; V1 builds no secondary pool.
        The disk tier is V2 only and has no auto policy, so dropping its budget
        leaves the tier out.
        """
        if kv_cache_config is None:
            return kv_cache_config
        dropped = {
            attr: None
            for attr in cls._OFFLOAD_TIER_BUDGET_ATTRS
            if getattr(kv_cache_config, attr)
        }
        if not dropped:
            return kv_cache_config
        return kv_cache_config.model_copy(update=dropped)

    def build_managers(self,
                       resources: Dict,
                       estimating_kv_cache: bool = False) -> None:
        """Construct KV caches for model and draft model (if applicable)."""
        if self._skip_est:
            self.configure_kv_cache_capacity()
        original_max_seq_len = self._max_seq_len

        # For encoder-decoder models, split the self/cross budgets first so
        # every enc-dec build creates a real cross pool.  This must happen
        # before any draft split so that the draft split operates on the
        # already-reduced self-pool budget.
        self_kv_cache_config = self._kv_cache_config
        cross_kv_cache_config = None
        if self._is_encoder_decoder():
            self_kv_cache_config, cross_kv_cache_config = self._split_kv_cache_budget_for_cross(
            )

        has_separate_one_model_draft = (
            self._draft_model_engine is None
            and self._should_create_separate_draft_kv_cache())
        # One term per axis: a draft pool exists to pair with, block reuse is on,
        # and the span is known. Anything outside keeps its unpaired path.
        self._joint_kv_cache_reuse = (has_separate_one_model_draft and
                                      self_kv_cache_config.enable_block_reuse
                                      and self._joint_reuse_supported())

        # Estimation managers are throwaway probes whose pools only hold dummy
        # requests, so an explicit offload tier would reserve capacity the probe
        # cannot fill. Encoder-decoder runs skip estimation, so dropping the
        # cross budgets is defensive.
        if estimating_kv_cache:
            self_kv_cache_config = self._drop_explicit_offload_tier_budgets(
                self_kv_cache_config)
            cross_kv_cache_config = self._drop_explicit_offload_tier_budgets(
                cross_kv_cache_config)

        # Split combined KV cache budgets before creating managers.
        has_draft = (
            self._draft_model_engine is not None  # two-model
            or has_separate_one_model_draft)  # one-model
        draft_kv_cache_config = None
        if has_draft:
            # The GPU split applies when each manager sizes its pools from
            # max_gpu_total_bytes (V2 and V1 VSWA). V1 non-VSWA and estimation
            # size GPU pools from a shared max_tokens instead.
            needs_gpu_split = (not estimating_kv_cache
                               and self._needs_gpu_kv_cache_budget_split(
                                   original_max_seq_len, self_kv_cache_config))
            if needs_gpu_split:
                self_kv_cache_config, draft_kv_cache_config = (
                    self._split_kv_cache_budget_for_draft(
                        "max_gpu_total_bytes", self_kv_cache_config,
                        draft_kv_cache_config))
            for budget_attr in self._OFFLOAD_TIER_BUDGET_ATTRS:
                self_kv_cache_config, draft_kv_cache_config = (
                    self._split_kv_cache_budget_for_draft(
                        budget_attr, self_kv_cache_config,
                        draft_kv_cache_config))

        compression_config = self._llm_args.kv_cache_compression_config
        compression_manager = create_kv_cache_compression_manager(
            compression_config,
            model_engine=self._model_engine,
            kv_cache_config=self_kv_cache_config,
            estimating_kv_cache=estimating_kv_cache and not self._skip_est,
        )
        cold_page_codec_provider = (
            compression_manager if compression_manager is not None
            and compression_manager.provides_cold_page_codec else None)

        kv_cache_manager = self._create_kv_cache_manager(
            self._model_engine,
            estimating_kv_cache,
            kv_cache_config_override=self_kv_cache_config,
            cold_page_codec_provider=cold_page_codec_provider,
        )

        # Carry the fp8 context-MLA workspace admission cap (computed in configure_kv_cache_capacity) onto
        # the real KV manager so the scheduler reads it directly instead of re-deriving from pool layout.
        # The estimation build reserves nothing and runs throwaway fresh-prefill dummies, so leave the
        # attribute unset there (PyExecutor._get_ctx_mla_kv_len_cap does not cap during warmup).
        if not estimating_kv_cache and kv_cache_manager is not None:
            kv_cache_manager.fp8_ctx_mla_kv_len_cap = self._fp8_ctx_mla_kv_len_cap

        if (not estimating_kv_cache and self._kv_connector_manager is not None
                and self._draft_model_engine is not None):
            raise NotImplementedError(
                "Connector manager is not supported for draft model.")

        draft_kv_cache_manager = None
        draft_build_kv_cache_config = (draft_kv_cache_config
                                       if draft_kv_cache_config is not None else
                                       self_kv_cache_config)

        # Two-model speculative decoding: draft model has separate engine
        if self._draft_model_engine is not None:
            if (self._is_kv_cache_manager_v2
                    and draft_kv_cache_config is not None):
                # Offload budgets are divided per manager, GPU budgets are not.
                assert (draft_kv_cache_config.max_gpu_total_bytes ==
                        self_kv_cache_config.max_gpu_total_bytes), (
                            "KVCacheManagerV2 does not support two-model "
                            "speculative decoding with separate draft GPU "
                            "budgets.")
            draft_kv_cache_manager = self._create_kv_cache_manager(
                self._draft_model_engine,
                estimating_kv_cache,
                kv_cache_config_override=draft_build_kv_cache_config)
        # One-model speculative decoding with different KV layouts
        elif self._should_create_separate_draft_kv_cache():
            draft_kv_cache_manager = self._create_one_model_draft_kv_cache_manager(
                original_max_seq_len,
                estimating_kv_cache,
                kv_cache_config_override=draft_build_kv_cache_config,
                cold_page_codec_provider=cold_page_codec_provider)

        # Encoder-decoder cross-attention pool
        cross_kv_cache_manager = None
        if cross_kv_cache_config is not None:
            cross_kv_cache_manager = self._create_cross_kv_cache_manager(
                cross_kv_cache_config, estimating_kv_cache,
                original_max_seq_len)

        resources[ResourceManagerType.KV_CACHE_MANAGER] = kv_cache_manager
        resources[
            ResourceManagerType.DRAFT_KV_CACHE_MANAGER] = draft_kv_cache_manager
        resources[
            ResourceManagerType.CROSS_KV_CACHE_MANAGER] = cross_kv_cache_manager
        if (compression_manager is not None
                and compression_manager.uses_iteration_lifecycle):
            resources[ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER] = (
                compression_manager)

    def teardown_managers(self, resources: Dict) -> None:
        """Clean up KV caches for model, draft model, and cross pool."""
        compression_manager = resources.pop(
            ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER, None)
        if compression_manager is not None:
            compression_manager.shutdown()
        resources[ResourceManagerType.KV_CACHE_MANAGER].shutdown()
        del resources[ResourceManagerType.KV_CACHE_MANAGER]
        draft_kv_cache_manager = resources[
            ResourceManagerType.DRAFT_KV_CACHE_MANAGER]
        if draft_kv_cache_manager:
            draft_kv_cache_manager.shutdown()
        del resources[ResourceManagerType.DRAFT_KV_CACHE_MANAGER]
        cross_kv_cache_manager = resources.get(
            ResourceManagerType.CROSS_KV_CACHE_MANAGER)
        if cross_kv_cache_manager is not None:
            cross_kv_cache_manager.shutdown()
        if ResourceManagerType.CROSS_KV_CACHE_MANAGER in resources:
            del resources[ResourceManagerType.CROSS_KV_CACHE_MANAGER]


def _build_per_layer_num_kv_heads(
    num_key_value_heads: int,
    num_hidden_layers: int,
    spec_config: Optional[SpeculativeConfig] = None,
    draft_config: Optional[ModelConfig] = None,
) -> Union[int, List[int]]:
    """
    Returns:
        An int when all layers share the same num_kv_heads (common case),
        or a list of num_kv_heads (one entry per layer) when target and
        draft models differ.
    """
    if spec_config is None or draft_config is None:
        return num_key_value_heads

    from ..speculative.utils import get_num_spec_layers
    draft_pretrained = draft_config.pretrained_config
    draft_num_kv_heads = getattr(
        draft_pretrained, 'num_key_value_heads',
        getattr(draft_pretrained, 'num_attention_heads', None))

    if draft_num_kv_heads is None or draft_num_kv_heads == num_key_value_heads:
        return num_key_value_heads

    num_spec_layers = get_num_spec_layers(spec_config)
    logger.info(f"Per-layer KV heads for speculative decoding: "
                f"target={num_key_value_heads} x {num_hidden_layers} layers, "
                f"draft={draft_num_kv_heads} x {num_spec_layers} layers, "
                f"total={num_hidden_layers + num_spec_layers} layers")
    return [num_key_value_heads] * num_hidden_layers + [draft_num_kv_heads
                                                        ] * num_spec_layers


def _get_mamba_cache_layer_masks(
    mamba_params: MambaKVCacheParams,
    mapping: Mapping,
    spec_config: Optional[SpeculativeConfig],
    is_draft: bool,
) -> tuple[List[bool], List[bool]]:
    use_separate_draft_kv_cache = (
        not mapping.enable_attention_dp
        and should_use_separate_draft_kv_cache(spec_config))
    return mamba_params.get_layer_masks(
        is_draft=is_draft,
        use_separate_draft_kv_cache=use_separate_draft_kv_cache,
    )


# The V1 hybrid managers select the convolution-state layout by model_type;
# MambaHybridCacheManagerV2 takes the layout by name and rejects model_type.
_CONV_STATE_LAYOUT_BY_MODEL_TYPE = {
    "nemotron_hybrid": "x_b_c",
    "qwen3_next": "q_k_v",
}


def _mamba_conv_layout_kwargs(kv_cache_manager_cls: type,
                              model_type: str) -> dict:
    """Constructor kwarg selecting the conv-state layout for a hybrid manager.

    Keeps the V1-vs-V2 dispatch in one place: a manager branch that forgets it
    would previously get V2's silent "x_b_c" default (the Kimi K3 bug fixed in
    this change).
    """
    if issubclass(kv_cache_manager_cls, MambaHybridCacheManagerV2):
        return {
            "conv_state_layout": _CONV_STATE_LAYOUT_BY_MODEL_TYPE[model_type]
        }
    return {"model_type": model_type}


def _get_qwen4_exp_ple_cache_params(config, *, total_layers: int,
                                    is_draft: bool):
    """Align target-only PLE state with a target/draft cache layout."""
    if is_draft:
        return None

    params = extract_qwen4_exp_ple_cache_params(config)
    num_target_layers = len(params.ple_layer_mask)
    if num_target_layers > total_layers:
        raise ValueError(
            "PLE layer mask cannot exceed the hybrid cache layout: "
            f"got {num_target_layers}, expected at most {total_layers}")
    if num_target_layers == total_layers:
        return params

    # Unified one-model caches append attention-only MTP layers.
    return dataclasses.replace(
        params,
        ple_layer_mask=params.ple_layer_mask + [False] *
        (total_layers - num_target_layers),
    )


def _create_kv_cache_manager(
        model_engine: Optional[PyTorchModelEngine],
        kv_cache_manager_cls,
        mapping: Mapping,
        kv_cache_config: KvCacheConfig,
        tokens_per_block: int,
        max_seq_len: int,
        max_batch_size: int,
        spec_config: Optional[SpeculativeConfig],
        sparse_attention_config: Optional[SparseAttentionConfig],
        max_num_tokens: int,
        max_beam_width: int,
        kv_connector_manager: Optional[KvCacheConnectorManager],
        estimating_kv_cache: bool = False,
        enable_kv_cache_stats: bool = False,
        execution_stream: Optional[torch.cuda.Stream] = None,
        # Optional overrides for one-model draft case (when model_engine is None)
        model_config: Optional[ModelConfig] = None,
        dtype: Optional[torch.dtype] = None,
        is_draft: Optional[bool] = None,
        layer_mask: Optional[List[bool]] = None,
        num_layers: Optional[int] = None,
        num_kv_heads: Optional[Union[int, List[int]]] = None,
        head_dim: Optional[int] = None,
        kv_cache_type=None,
        is_disagg: bool = False,
        disable_overlap_scheduler: bool = False,
        cold_page_codec_provider: Optional[object] = None,
        kv_events_config: Optional[KVEventsConfig] = None,
        joint_kv_cache_reuse: bool = False,
        max_cuda_graph_batch_size: Optional[int] = None) -> KVCacheManager:
    """
    Returns:
        A KVCacheManager instance for the given model engine or model config
    """
    if cold_page_codec_provider is not None and not issubclass(
            kv_cache_manager_cls, KVCacheManagerV2):
        raise ValueError(
            "Cold-page quantization requires the resolved KV cache manager "
            f"to be KVCacheManagerV2; selected {kv_cache_manager_cls.__name__}")

    if (estimating_kv_cache
            and issubclass(kv_cache_manager_cls, KVCacheManagerV2)
            and kv_cache_config.pool_ratio is None
            and kv_cache_config.avg_seq_len is not None
            and kv_cache_config.avg_seq_len > max_seq_len):
        # Estimation can build multiple managers from the same temporary
        # config. The first manager may reduce max_seq_len to fit max_tokens,
        # so later draft/cross managers need a per-manager workload length.
        # Keep the shared config untouched because it is restored after
        # estimation.
        kv_cache_config = kv_cache_config.model_copy(
            update={"avg_seq_len": max_seq_len})

    # Extract config from model_engine or use provided model_config
    if model_config is not None:
        config = model_config.pretrained_config
        quant_config = model_config.quant_config
        _model_config = model_config
    else:
        config = model_engine.model.model_config.pretrained_config
        quant_config = model_engine.model.model_config.quant_config
        _model_config = model_engine.model.model_config

    if dtype is None:
        dtype = model_engine.dtype

    if is_draft is None:
        is_draft = model_engine.is_draft_model

    if kv_cache_type is None:
        kv_cache_type = tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF

    hidden_size = config.hidden_size
    num_attention_heads = config.num_attention_heads
    num_key_value_heads = num_kv_heads if num_kv_heads is not None else getattr(
        config, 'num_key_value_heads', num_attention_heads)
    if not isinstance(head_dim, int):
        head_dim = getattr(config, "head_dim", None)
    if not isinstance(head_dim, int):
        head_dim = hidden_size // num_attention_heads

    # Gemma4: build per-layer head_dim, num_kv_heads, and sliding window
    # for hybrid attention. Different layer types need different KV cache
    # pool groups (via max_attention_window) so FlashInfer page indices
    # are consistent within each group.
    if is_gemma4_hybrid(config):
        layer_types = config.layer_types
        global_head_dim = config.global_head_dim
        attention_k_eq_v = getattr(config, 'attention_k_eq_v', False)
        num_global_kv_heads = (getattr(config, 'num_global_key_value_heads',
                                       None) or num_key_value_heads)
        head_dim_list = []
        kv_heads_list = []
        for lt in layer_types:
            is_sliding = (lt == "sliding_attention")
            if is_sliding:
                head_dim_list.append(head_dim)
                kv_heads_list.append(num_key_value_heads)
            else:
                head_dim_list.append(global_head_dim)
                use_k_eq_v = attention_k_eq_v and not is_sliding
                kv_heads_list.append(
                    num_global_kv_heads if use_k_eq_v else num_key_value_heads)
        head_dim = head_dim_list
        num_key_value_heads = kv_heads_list

    # Derive per-layer max_attention_window for any model that publishes a
    # mixed sliding/full `layer_types` schedule (Gemma4 hybrid included) so V2
    # creates separate pool groups for sliding vs full-attention layers;
    # otherwise every layer lands in one full-context pool and the bounded
    # layers keep blocks they can never read.  Sliding layers get the model's
    # `sliding_window`, full layers `max_seq_len`; V2 then evicts old blocks
    # once kv_len exceeds the window (only ~ceil(sliding_window/page_size)
    # blocks per sequence for sliding layers), and FlashInfer's prepare()
    # picks up the smaller per-pool block count automatically.  Gemma4 hybrid
    # always resolves to KVCacheManagerV2 (see _non_hybrid_kv_cache_manager_cls),
    # so the V2 guard never excludes it.  A user-supplied max_attention_window
    # always wins.  `_derive_v2_layer_type_attention_windows` applies the same
    # rules in the creator's static cost model, so the budget split sizes the
    # manager from the windows it is built with; it also keeps the default
    # when a configured `pool_ratio` does not match the derived layer groups.
    # Skip derivation for the one-model draft manager: (1) in one-model
    # spec-decode the KV memory budget is already split between target and
    # draft, so VSWA sizing here would size each window pool from the unsplit
    # free-memory budget; (2) draft layers live at global indices past the
    # target's num_hidden_layers, so _project_max_attention_window_vec would
    # wrap them back onto pattern[0]. The draft config sets
    # max_attention_window=None to opt out, not to request derivation.
    # The cross-attention pool holds encoder-side KV that the decoder's
    # `layer_types` do not describe, so it keeps the default too.
    # Keep estimation storage full-context unless distinct page layouts require
    # windowed storage (Gemma4). Enable automatic windowing for the final manager.
    derived_windows = None
    if (not is_draft and kv_cache_type
            == tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF
            and (not estimating_kv_cache or is_gemma4_hybrid(config))):
        derived_windows = _derive_v2_layer_type_attention_windows(
            kv_cache_config, kv_cache_manager_cls, _model_config, max_seq_len)
    if derived_windows is not None:
        if layer_mask is None and spec_config is not None:
            # `get_pp_layers` appends the one-model speculative layers after
            # the decoder layers when they share this manager, and V2 resolves
            # each layer's window by global layer id: give them the full
            # context the single-window default gave them instead of wrapping
            # them onto the decoder pattern.
            derived_windows = derived_windows + [int(
                max_seq_len)] * get_num_spec_layers(spec_config)
        assert uses_vswa_kv_cache_layout(derived_windows), (
            "derived per-layer windows must select a VSWA layout; a non-VSWA "
            "vector would reshape the single pool instead of splitting it")
        logger.info(
            "Derived per-layer max_attention_window from layer_types for "
            f"{kv_cache_manager_cls.__name__}: {derived_windows} "
            f"({len(set(derived_windows))} distinct windows)")
        kv_cache_config = copy.copy(kv_cache_config)
        kv_cache_config.max_attention_window = derived_windows

    # Note: Gemma4 KV sharing is handled at the model level — shared layers
    # use cache_layer_idx to read from the target layer's cache slot via
    # Gemma4Attention. No layer_mask exclusion needed here.

    if quant_config is not None and quant_config.quant_mode.has_fp8_kv_cache():
        kv_cache_dtype = tensorrt_llm.bindings.DataType.FP8
    elif quant_config is not None and quant_config.quant_mode.has_fp4_kv_cache(
    ):
        kv_cache_dtype = tensorrt_llm.bindings.DataType.NVFP4
    else:
        kv_cache_dtype = str_dtype_to_binding(torch_dtype_to_str(dtype))

    # Use provided num_layers if available, otherwise use config.
    # When layer_mask is set (e.g., KV sharing), num_layers for the cache
    # manager must equal the number of enabled (True) layers in the mask.
    if num_layers is not None:
        num_hidden_layers = num_layers
    elif layer_mask is not None:
        num_hidden_layers = sum(layer_mask)
    else:
        num_hidden_layers = config.num_hidden_layers
    # Only include draft KV heads in the per-layer list when draft layers
    # are NOT handled by a separate draft KV cache manager.  When layer_mask
    # is provided from the caller, it means the main KV cache covers only
    # the masked (target) layers and draft layers live in their own manager.
    draft_config_for_kv = None
    if layer_mask is None:
        draft_config_for_kv = (getattr(model_engine.model, 'draft_config', None)
                               if model_engine is not None else None)
    # If num_key_value_heads is already a per-layer list (e.g., Gemma4 hybrid),
    # use it directly; otherwise build from the scalar value.
    if isinstance(num_key_value_heads, list):
        per_layer_num_kv_heads = num_key_value_heads
    else:
        per_layer_num_kv_heads = _build_per_layer_num_kv_heads(
            num_key_value_heads, num_hidden_layers, spec_config,
            draft_config_for_kv)
    manager_extra_kwargs = {}
    if issubclass(kv_cache_manager_cls, KVCacheManagerV2):
        manager_extra_kwargs["max_cuda_graph_batch_size"] = (
            model_engine._max_cuda_graph_batch_size
            if model_engine is not None else max_cuda_graph_batch_size)
        manager_extra_kwargs["enable_stats"] = enable_kv_cache_stats
        manager_extra_kwargs[
            "cold_page_codec_provider"] = cold_page_codec_provider
        manager_extra_kwargs["kv_events_config"] = kv_events_config
        manager_extra_kwargs["joint_kv_cache_reuse"] = joint_kv_cache_reuse
        manager_extra_kwargs[
            "disable_overlap_scheduler"] = disable_overlap_scheduler
        # Vocab size also enables multimodal event decoding and its per-block
        # digest scan. Leave it unset for text-only engines. Without an engine
        # (e.g. separate one-model draft caches), preserve config resolution:
        # a text sub-config alone cannot identify a multimodal deployment.
        needs_multimodal_keys = model_engine is None or model_engine.is_multimodal
        manager_extra_kwargs["vocab_size"] = (resolve_vocab_size(config) if
                                              needs_multimodal_keys else None)
        if (needs_multimodal_keys and manager_extra_kwargs["vocab_size"] is None
                and kv_cache_config.enable_block_reuse):
            logger.warning(
                "Could not resolve vocab_size from the model config; "
                "multimodal requests will fail when block reuse is on. "
                "Disable block reuse to serve them.")
    elif kv_events_config is not None and kv_events_config.enable_kv_cache_events:
        logger.warning(
            "kv_cache_config.kv_events_config is set but streaming KV event "
            "publishing requires KV cache manager V2; events will not be "
            f"published for {kv_cache_manager_cls.__name__}.")
    if issubclass(kv_cache_manager_cls, MambaHybridCacheManagerV2):
        manager_extra_kwargs["is_disagg"] = is_disagg

    if is_kimi_linear(config):
        # Kimi K3 hybrid: KDA (Kimi Delta Attention) recurrent/conv states on
        # the mamba side of the hybrid manager, absorbed-MQA MLA latent cache
        # (num_kv_heads=1, head_dim = kv_lora_rank + qk_rope_head_dim,
        # SELFKONLY) on the paged-KV side. Must come before the is_mla(...)
        # route: the kimi_linear config carries MLA fields, but only 24 of
        # its 93 layers are MLA.
        if max_beam_width > 1:
            raise ValueError(
                "MambaHybridCacheManager + beam search is not supported yet.")
        if not estimating_kv_cache and kv_connector_manager is not None:
            raise NotImplementedError(
                "Connector manager is not supported for MambaHybridCacheManager."
            )
        mamba_params = extract_mamba_kv_cache_params(
            config,
            spec_config=spec_config,
            quant_config=quant_config,
        )
        mamba_layer_mask, full_attention_layer_mask = (
            _get_mamba_cache_layer_masks(
                mamba_params,
                mapping,
                spec_config,
                is_draft,
            ))
        num_mamba_layers = (0 if is_draft and mamba_params.num_draft_layers > 0
                            else mamba_params.num_mamba_layers)
        # Kimi K3 KDA state sharding follows the attention-family TP
        # semantics (Qwen3-Next pattern): replicated under attention-DP,
        # head-sharded across tp_size otherwise. That is exactly the cache
        # manager's own internal gate (`tp_size = 1 if enable_attention_dp
        # else tp_size`, then num_heads / n_groups / conv_dim divide by
        # it), so the params pass through unscaled.
        # KDA fused multi-token verify (trtllm::kda_mtp_decode): when the
        # kernel can run here, allocate the per-slot replay caches instead
        # of the legacy per-step intermediate verification buffers. The
        # kernel replays accepted drafts from these caches and commits
        # states in place, replacing the intermediate-buffer + promotion
        # flow for KDA layers.
        kimi_extra_kwargs = {}
        kda_replay_manager_types = (MixedMambaHybridCacheManager,
                                    MambaHybridCacheManagerV2)
        if (spec_config is not None
                and issubclass(kv_cache_manager_cls, kda_replay_manager_types)):
            from ..modules.kimi_kda._kda_kernels import \
                is_kda_mtp_verify_available
            if is_kda_mtp_verify_available():
                kimi_extra_kwargs["kda_replay_num_spec"] = (
                    spec_config.tokens_per_gen_step - 1)
        # KDA's conv state is a [Q | K | V] concatenation whose three sections
        # have identical width, i.e. the qwen3_next section layout.
        kimi_extra_kwargs.update(
            _mamba_conv_layout_kwargs(kv_cache_manager_cls, "qwen3_next"))
        kv_cache_manager = kv_cache_manager_cls(
            # mamba (KDA) cache parameters
            mamba_params.state_size,
            mamba_params.conv_kernel,
            mamba_params.num_heads,
            mamba_params.n_groups,
            mamba_params.head_dim,
            num_mamba_layers,
            mamba_layer_mask,
            mamba_params.dtype,
            mamba_params.mamba_ssm_cache_dtype,
            # kv cache parameters (MLA latent cache)
            kv_cache_config,
            tensorrt_llm.bindings.internal.batch_manager.CacheType.SELFKONLY,
            num_layers=sum(full_attention_layer_mask),
            layer_mask=full_attention_layer_mask,
            num_kv_heads=1,
            head_dim=config.kv_lora_rank + config.qk_rope_head_dim,
            tokens_per_block=tokens_per_block,
            max_seq_len=max_seq_len,
            max_num_tokens=max_num_tokens,
            is_draft=is_draft,
            max_batch_size=max_batch_size,
            mapping=mapping,
            dtype=kv_cache_dtype,
            spec_config=spec_config,
            is_estimating_kv_cache=estimating_kv_cache,
            execution_stream=execution_stream,
            **kimi_extra_kwargs,
            **manager_extra_kwargs,
        )
    elif is_mla(config):
        kv_cache_manager = kv_cache_manager_cls(
            kv_cache_config,
            tensorrt_llm.bindings.internal.batch_manager.CacheType.SELFKONLY,
            num_layers=num_hidden_layers,
            num_kv_heads=1,
            head_dim=config.kv_lora_rank + config.qk_rope_head_dim,
            tokens_per_block=tokens_per_block,
            max_seq_len=max_seq_len,
            max_batch_size=max_batch_size,
            mapping=mapping,
            dtype=kv_cache_dtype,
            spec_config=spec_config,
            max_num_tokens=max_num_tokens,
            max_beam_width=max_beam_width,
            is_draft=is_draft,
            kv_connector_manager=kv_connector_manager
            if not estimating_kv_cache else None,
            sparse_attention_config=sparse_attention_config,
            pretrained_config=config,
            is_estimating_kv_cache=estimating_kv_cache,
            execution_stream=execution_stream,
            layer_mask=layer_mask,
            is_disagg=is_disagg,
            **manager_extra_kwargs,
        )
    elif is_nemotron_hybrid(config):
        if max_beam_width > 1:
            raise ValueError(
                "MambaHybridCacheManager + beam search is not supported yet.")

        if not estimating_kv_cache and kv_connector_manager is not None:
            raise NotImplementedError(
                "Connector manager is not supported for MambaHybridCacheManager."
            )

        mamba_params = extract_mamba_kv_cache_params(
            config,
            spec_config=spec_config,
            quant_config=quant_config,
        )
        mamba_layer_mask, full_attention_layer_mask = (
            _get_mamba_cache_layer_masks(
                mamba_params,
                mapping,
                spec_config,
                is_draft,
            ))
        num_mamba_layers = (0 if is_draft and mamba_params.num_draft_layers > 0
                            else mamba_params.num_mamba_layers)

        # Replay state update kernel for MTP: default on for sm >= 80; gates
        # below disable it for incompatible feature combinations.  Cpp cache
        # manager doesn't expose use_replay_state_update, so the wrapper
        # property's getattr default keeps replay off there automatically.
        sm = get_sm_version()
        stochastic_rounding = getattr(
            quant_config, 'mamba_ssm_stochastic_rounding',
            False) if quant_config is not None else False

        use_replay = spec_config is not None and sm >= 80
        if spec_config is None:
            logger.info(
                "Replay kernel requires speculative decoding; using non-replay path"
            )

        # Block reuse (prefix caching): replay leaves SSM state at a
        # checkpoint after speculation. The next decode step replays forward
        # to correct it. If block reuse feeds that stale state into a new
        # prefill, the correction never happens.
        # Currently we only save and reuse context tokens so this does not affect.

        # Tree attention: replay assumes linear token sequence.
        if (spec_config is not None
                and getattr(spec_config, 'use_dynamic_tree', False)):
            logger.info("Replay kernel incompatible with tree attention; "
                        "using legacy MTP path")
            use_replay = False

        # Replay Philox uses PTX cvt.rs.f16x2.f32 which needs 100 <= sm < 120.
        # Flashinfer has a SW fallback at any SM.
        if (stochastic_rounding
                and mamba_params.mamba_ssm_cache_dtype == torch.float16
                and (sm < 100 or sm in (120, 121))):
            logger.info("Replay kernel Philox requires 100 <= sm < 120; "
                        "using legacy MTP path for stochastic rounding support")
            use_replay = False

        # Use replay algorithm for mamba (default is on).
        enforce_disable_replay = os.environ.get('TRTLLM_USE_MAMBA_REPLAY',
                                                '1') == '0'
        if enforce_disable_replay:
            logger.info(
                "Replay kernel is disabled by TRTLLM_USE_MAMBA_REPLAY=0")
            use_replay = False
        else:
            logger.info(
                "Replay kernel is not changed since TRTLLM_USE_MAMBA_REPLAY=1")

        # Stochastic-rounding seeds must live on the cache manager (not be
        # re-created with torch.randint per forward) whenever SR can fire
        # on the fp16 SSM cache.  This mirrors the predicate the mixer uses
        # internally (`_stochastic_rounding_for_flashinfer` /
        # `_stochastic_rounding_for_replay`) so allocation matches consumption.
        mamba_ssm_stochastic_rounding = (stochastic_rounding
                                         and mamba_params.mamba_ssm_cache_dtype
                                         == torch.float16)
        mamba_manager_extra_kwargs = dict(manager_extra_kwargs)
        mamba_manager_extra_kwargs.update(
            _mamba_conv_layout_kwargs(kv_cache_manager_cls, "nemotron_hybrid"))
        kv_cache_manager = kv_cache_manager_cls(
            # mamba cache parameters
            mamba_params.state_size,
            mamba_params.conv_kernel,
            mamba_params.num_heads,
            mamba_params.n_groups,
            mamba_params.head_dim,
            num_mamba_layers,
            mamba_layer_mask,
            mamba_params.dtype,
            mamba_params.mamba_ssm_cache_dtype,
            # kv cache parameters
            kv_cache_config,
            tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF,
            num_layers=sum(full_attention_layer_mask),
            layer_mask=full_attention_layer_mask,
            num_kv_heads=per_layer_num_kv_heads,
            head_dim=head_dim,
            tokens_per_block=tokens_per_block,
            max_seq_len=max_seq_len,
            max_num_tokens=max_num_tokens,
            is_draft=is_draft,
            max_batch_size=max_batch_size,
            mapping=mapping,
            dtype=kv_cache_dtype,
            spec_config=spec_config,
            is_estimating_kv_cache=estimating_kv_cache,
            execution_stream=execution_stream,
            use_replay_state_update=use_replay,
            mamba_ssm_stochastic_rounding=mamba_ssm_stochastic_rounding,
            **mamba_manager_extra_kwargs,
        )
    elif is_qwen3_hybrid(config) or is_qwen4_exp(config):
        if max_beam_width > 1:
            raise ValueError(
                "MambaHybridCacheManager + beam search is not supported yet.")

        if not estimating_kv_cache and kv_connector_manager is not None:
            raise NotImplementedError(
                "Connector manager is not supported for MambaHybridCacheManager."
            )
        mamba_params = extract_mamba_kv_cache_params(
            config,
            spec_config=spec_config,
            quant_config=quant_config,
        )
        mamba_layer_mask, full_attention_layer_mask = (
            _get_mamba_cache_layer_masks(
                mamba_params,
                mapping,
                spec_config,
                is_draft,
            ))
        num_mamba_layers = (0 if is_draft and mamba_params.num_draft_layers > 0
                            else mamba_params.num_mamba_layers)
        # Replay state update for GDN MTP: mirrors the nemotron_hybrid gating
        # above, minus the Mamba2-specific stochastic-rounding/Philox gate.
        # The GDN replay kernel does a plain cast on checkpoint commit, so
        # quantized SSM cache dtypes stay on the legacy path.
        sm = get_sm_version()
        use_replay = spec_config is not None and sm >= 80
        if spec_config is None:
            logger.info(
                "GDN replay kernel requires speculative decoding; using "
                "non-replay path")
        elif spec_config.tokens_per_gen_step > 8:
            logger.info("GDN cached replay supports at most 8 tokens per "
                        "generation step; using non-replay path")
            use_replay = False

        # Tree attention: replay assumes a linear token sequence.
        if (spec_config is not None
                and getattr(spec_config, 'use_dynamic_tree', False)):
            logger.info("GDN replay kernel incompatible with tree attention; "
                        "using legacy MTP path")
            use_replay = False

        if mamba_params.mamba_ssm_cache_dtype not in (torch.float32,
                                                      torch.bfloat16,
                                                      torch.float16):
            logger.info(
                "GDN replay kernel does not support quantized SSM cache "
                f"dtype {mamba_params.mamba_ssm_cache_dtype}; using legacy "
                "MTP path")
            use_replay = False

        # Replay is enabled by default for eligible GDN MTP workloads.
        if not is_gdn_replay_enabled():
            use_replay = False

        # GDN replay supports the contiguous C++ V1 state pool and the indirect
        # per-layer state views exposed by V2. Mixed/Python does not expose an
        # all-layer commit, so keep that manager but use non-replay MTP.
        replay_manager_types = (CppMambaHybridCacheManager,
                                MambaHybridCacheManagerV2)
        if use_replay and not issubclass(kv_cache_manager_cls,
                                         replay_manager_types):
            logger.info("GDN replay requires C++ V1 or V2 Mamba cache manager; "
                        f"{kv_cache_manager_cls.__name__} was selected, so the "
                        "non-replay MTP path will be used")
            use_replay = False
        logger.info("GDN replay state update: " +
                    ("ENABLED" if use_replay else "DISABLED"))

        mamba_manager_extra_kwargs = dict(manager_extra_kwargs)
        mamba_manager_extra_kwargs.update(
            _mamba_conv_layout_kwargs(kv_cache_manager_cls, "qwen3_next"))
        if getattr(sparse_attention_config, "algorithm", None) == "qsa":
            # Resolve the side-cache shape from the same checkpoint geometry
            # used to construct the QSA index projection.
            mamba_manager_extra_kwargs.update(
                sparse_attention_config=sparse_attention_config,
                pretrained_config=config,
            )
        if is_qwen4_exp(config) and issubclass(kv_cache_manager_cls,
                                               MambaHybridCacheManagerV2):
            ple_cache_params = _get_qwen4_exp_ple_cache_params(
                config,
                total_layers=len(mamba_layer_mask),
                is_draft=is_draft,
            )
            if ple_cache_params is not None:
                mamba_manager_extra_kwargs[
                    "qwen4_exp_ple_cache_params"] = ple_cache_params
        kv_cache_manager = kv_cache_manager_cls(
            # mamba cache parameters
            mamba_params.state_size,
            mamba_params.conv_kernel,
            mamba_params.num_heads,
            mamba_params.n_groups,
            mamba_params.head_dim,
            num_mamba_layers,
            mamba_layer_mask,
            mamba_params.dtype,
            mamba_params.mamba_ssm_cache_dtype,
            # kv cache parameters
            kv_cache_config,
            tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF,
            num_layers=sum(full_attention_layer_mask),
            layer_mask=full_attention_layer_mask,
            num_kv_heads=per_layer_num_kv_heads,
            head_dim=head_dim,
            tokens_per_block=tokens_per_block,
            max_seq_len=max_seq_len,
            max_num_tokens=max_num_tokens,
            is_draft=is_draft,
            max_batch_size=max_batch_size,
            mapping=mapping,
            dtype=kv_cache_dtype,
            spec_config=spec_config,
            is_estimating_kv_cache=estimating_kv_cache,
            execution_stream=execution_stream,
            use_replay_state_update=use_replay,
            **mamba_manager_extra_kwargs,
        )
    else:
        # NOTE: this is a workaround for VSWA to switch to calculate_max_num_blocks_for_vswa in KVCahceManager
        # Only needed for V1; V2 handles per-layer windows natively via life cycles.
        is_vswa = uses_vswa_kv_cache_layout(
            kv_cache_config.max_attention_window)
        binding_model_config = None
        if is_vswa and kv_cache_manager_cls.__name__ == "KVCacheManager":
            binding_model_config = _model_config.get_bindings_model_config(
                tokens_per_block=tokens_per_block,
                kv_cache_config=kv_cache_config,
                spec_config=spec_config)

        # KVCacheManager (V1) doesn't support per-layer head_dim lists;
        # use max for estimation. KVCacheManagerV2 handles lists natively.
        effective_head_dim = (
            max(head_dim) if isinstance(head_dim, list)
            and kv_cache_manager_cls.__name__ == "KVCacheManager" else head_dim)
        kv_cache_manager = kv_cache_manager_cls(
            kv_cache_config,
            kv_cache_type,
            num_layers=num_hidden_layers,
            num_kv_heads=per_layer_num_kv_heads,
            head_dim=effective_head_dim,
            tokens_per_block=tokens_per_block,
            max_seq_len=max_seq_len,
            max_batch_size=max_batch_size,
            mapping=mapping,
            dtype=kv_cache_dtype,
            spec_config=spec_config,
            max_num_tokens=max_num_tokens,
            model_config=binding_model_config,
            max_beam_width=max_beam_width,
            is_draft=is_draft,
            kv_connector_manager=kv_connector_manager
            if not estimating_kv_cache else None,
            sparse_attention_config=sparse_attention_config,
            pretrained_config=config,
            is_estimating_kv_cache=estimating_kv_cache,
            execution_stream=execution_stream,
            layer_mask=layer_mask,
            is_disagg=is_disagg,
            **manager_extra_kwargs,
        )
    # Note: Gemma4 KV sharing cache remapping is handled in Gemma4Attention
    # via cache_layer_idx — shared layers use target layer's index for
    # get_buffers(). No layer_offsets remapping needed here.

    # Propagate the finalized chunked-prefill flag so KVCacheManager.fit_token_budget
    # only shrinks context chunks when the attention backend can consume a
    # partial context chunk. The flag is read from attn_runtime_features, which
    # py_executor_creator finalizes (including the SM-version /
    # attention-backend overrides) before build_managers runs.
    if isinstance(kv_cache_manager,
                  KVCacheManager) and model_engine is not None:
        kv_cache_manager.enable_chunked_prefill = bool(
            model_engine.attn_runtime_features.chunked_prefill)

    return kv_cache_manager


def validate_kv_cache_compression_compatibility(
    config: KvCacheCompressionConfig,
    kv_cache_config: KvCacheConfig,
    spec_config: Optional[SpeculativeConfig],
) -> None:
    """Reject unsupported KV-cache compression feature combinations."""
    if config.algorithm == "quantization_for_cold_page":
        from tensorrt_llm.runtime.kv_cache_manager_v2 import _BACKEND

        if _BACKEND == "python":
            raise ValueError(
                "Cold-page quantization requires the C++ KVCacheManagerV2 backend"
            )
        if config.quant == "nvfp4" and not is_sm_100f():
            raise RuntimeError(
                "NVFP4 cold-page quantization requires an SM100-family device "
                "(SM100, SM103 or SM107).")
    elif config.algorithm == "triattention" and not is_sm_100f():
        raise RuntimeError(
            "TriAttention requires an SM100-family device (SM100 or SM103).")

    if kv_cache_config.enable_block_reuse and not config.supports_block_reuse():
        raise ValueError(
            f"KV-cache compression algorithm {config.algorithm!r} does not "
            "support KV-cache block reuse. Set "
            "KvCacheConfig.enable_block_reuse=False.")
    if spec_config is None:
        return
    if not config.supports_speculative_decoding():
        guidance = ("; TriAttention requires eviction_mode='union'"
                    if config.algorithm == "triattention" else "")
        raise ValueError(
            f"KV-cache compression algorithm {config.algorithm!r} does not "
            "support speculative decoding with its current configuration"
            f"{guidance}")
    mode = spec_config.spec_dec_mode
    if config.algorithm == "quantization_for_cold_page":
        supported = mode.is_mtp_eagle_one_model() or mode.is_eagle3_one_model()
        guidance = "one-model MTP-EAGLE or EAGLE3"
    else:
        supported = mode.is_mtp_one_model() or mode.is_eagle3_one_model()
        guidance = "one-model MTP or EAGLE3"
    if not supported:
        raise ValueError(
            f"KV-cache compression does not support speculative decoding "
            f"mode {mode.name}; use {guidance}")


def create_kv_cache_compression_manager(
    config: Optional[KvCacheCompressionConfig],
    *,
    model_engine: PyTorchModelEngine,
    kv_cache_config: KvCacheConfig,
    estimating_kv_cache: bool = False,
) -> Optional[KVCacheCompressionManager]:
    """Validate, select, and construct the configured manager before KVCM."""
    if config is None:
        return None
    if model_engine.mapping.has_cp_helix():
        # TODO: Revisit after KVCC validates HELIX-sharded Page ownership and migration.
        raise ValueError(
            "KV-cache compression does not support HELIX context parallelism.")

    if config.algorithm == "quantization_for_cold_page":
        if config.quant != "nvfp4":
            raise NotImplementedError(
                f"Unsupported cold-page quantization format {config.quant!r}")
        if estimating_kv_cache:
            return None
        quant_config = model_engine.model.model_config.quant_config
        if (quant_config is not None and getattr(
                quant_config, "kv_cache_quant_algo", None) == QuantAlgo.NVFP4):
            logger.info(
                "Skipping cold-page NVFP4 quantization because the active KV "
                "cache already uses NVFP4; KVCM will migrate it losslessly.")
            return None

        validate_kv_cache_compression_compatibility(config, kv_cache_config,
                                                    model_engine.spec_config)
        from ..kv_cache_compression.quantization_for_cold_page.nvfp4_quantization import \
            Nvfp4ColdPageQuantizationCompression

        return Nvfp4ColdPageQuantizationCompression(
            config,
            pretrained_config=model_engine.model.model_config.pretrained_config,
        )

    if config.algorithm == "triattention":
        validate_kv_cache_compression_compatibility(config, kv_cache_config,
                                                    model_engine.spec_config)
        # TriAttention imports CuTe/CUTLASS; keep normal executor startup lazy.
        from ..kv_cache_compression.triattention.triattention import \
            TriAttentionCompressionManager

        return TriAttentionCompressionManager(
            config,
            pretrained_config=model_engine.model.model_config.pretrained_config,
        )

    logger.warning(
        "KV-cache compression algorithm '%s' is not registered; running without "
        "a compression manager.",
        config.algorithm,
    )
    return None


def compute_max_num_sequences(mapping: Mapping,
                              max_batch_size: int,
                              disable_overlap_scheduler: bool,
                              enable_overlap_headroom: bool = False) -> int:
    """Size the sequence-slot pool (and the sampler state it indexes).

    ``enable_overlap_headroom`` is intentionally opt-in; see
    ``should_enable_overlap_headroom``. Pipeline parallelism already sizes the
    pool by ``pp_size``.
    """
    if mapping.has_pp():
        num_micro_batches = mapping.pp_size
    else:
        num_micro_batches = (2 if enable_overlap_headroom
                             and not disable_overlap_scheduler else 1)
    return max_batch_size * num_micro_batches


def resolve_max_num_sequences(model_engine,
                              mapping: Mapping,
                              max_batch_size: int,
                              llm_args,
                              max_num_sequences: Optional[int] = None) -> int:
    """Resolve the seat-pool size, preferring an explicit value, then the
    engine's published pool, then a fresh ``compute_max_num_sequences``."""
    if max_num_sequences is not None:
        return max_num_sequences
    engine_seats = getattr(model_engine, "max_num_seq_slots", None)
    if engine_seats is not None:
        return engine_seats
    return compute_max_num_sequences(mapping,
                                     max_batch_size,
                                     llm_args.disable_overlap_scheduler,
                                     enable_overlap_headroom=getattr(
                                         model_engine,
                                         "_enable_overlap_headroom", False))


def should_enable_adp_dummy_fixes(mapping: Mapping) -> bool:
    """Enable transactional ADP dummy handling while PP remains follow-up."""
    return not mapping.has_pp()


_VALIDATED_OVERLAP_ADP_DUMMY_MODEL_TYPES = ("deepseek_v4", "qwen3_5_moe")


def should_enable_scheduler_aware_adp_dummy(
        model_type: Optional[str], mapping: Mapping,
        disable_overlap_scheduler: bool) -> bool:
    """Enable scheduler-aware padding for validated lifecycle configurations."""
    return (should_enable_adp_dummy_fixes(mapping)
            and (disable_overlap_scheduler
                 or model_type in _VALIDATED_OVERLAP_ADP_DUMMY_MODEL_TYPES))


def should_enable_non_overlap_adp_forward_intent(
        mapping: Mapping, disable_overlap_scheduler: bool) -> bool:
    """Enable fresh cross-rank dummy intent for the generic non-overlap path."""
    return (should_enable_adp_dummy_fixes(mapping)
            and disable_overlap_scheduler)


def should_enable_overlap_headroom(mapping: Mapping,
                                   disable_overlap_scheduler: bool,
                                   kv_cache_manager_is_v2: bool,
                                   is_hybrid: bool = False,
                                   has_mrope_delta_cache: bool = False) -> bool:
    """Gate the extra micro-batch of sequence slots.

    True only where a retiring request and the replacement that took its place
    can own a seat at the same time: attention DP, non-PP, overlap-on, V2 and
    non-hybrid.

    Widening the pool is only safe when every ``py_seq_slot``-indexed pool is
    sized from ``compute_max_num_sequences``. Two model families size one from
    something else instead, so they keep the single-micro-batch pool:

    * ``is_hybrid``: ``MambaCacheManager`` re-derives its own capacity as
      ``max_batch_size * pp_size``, which a doubled non-PP pool would exhaust.
    * ``has_mrope_delta_cache``: Qwen2/2.5-VL and Qwen3-VL hold
      ``max_num_tokens * pp_size + 1`` MRoPE deltas while indexing them by
      ``py_seq_slot``, relying on ``max_batch_size <= max_num_tokens`` to stay in
      bounds. The top entry is the reserved dummy slot, so a doubled pool first
      aliases the dummy -- silently giving padded requests a real request's
      delta -- and then indexes past the end.
    """
    if is_hybrid or has_mrope_delta_cache or not kv_cache_manager_is_v2:
        return False
    return (mapping.enable_attention_dp and not mapping.has_pp()
            and not disable_overlap_scheduler)


def validate_seq_slot_pool_covers_admission(max_num_sequences: int,
                                            kv_cache_manager) -> None:
    """Fail at startup if the KV index pool cannot cover the seat pool.

    The check is one-sided on purpose: an index pool narrower than the seat pool
    silently defers admitted requests, while a wider one is legitimate. Managers
    that do not publish an integer ``max_admissible_sequences`` are skipped.
    """
    admissible = getattr(kv_cache_manager, "max_admissible_sequences", None)
    if not isinstance(admissible, int):
        return
    if admissible >= max_num_sequences:
        return
    raise ValueError(
        f"{type(kv_cache_manager).__name__} can lease KV cache indices for "
        f"{admissible} concurrent sequences but the executor's sequence-slot "
        f"pool holds {max_num_sequences}: the index pool is smaller than the "
        "seat pool, so admitted requests would be silently deferred one at a "
        "time (nvbug 6627795). The seat pool must come from "
        "_util.compute_max_num_sequences and the index pool must cover it; a "
        "shortfall means one of them was re-derived from max_batch_size.")


def create_py_executor_instance(
    *,
    dist,
    resources,
    mapping,
    llm_args,
    ctx_chunk_config,
    model_engine,
    start_worker,
    sampler,
    drafter,
    guided_decoder: Optional[GuidedDecoder] = None,
    lora_config: Optional[LoraConfig] = None,
    garbage_collection_gen0_threshold: Optional[int] = None,
    kv_connector_manager: Optional[KvCacheConnectorManager] = None,
    resource_governor_queue=None,
    max_seq_len: Optional[int] = None,
    max_batch_size: Optional[int] = None,
    max_beam_width: Optional[int] = None,
    max_num_tokens: Optional[int] = None,
    peft_cache_config: Optional[PeftCacheConfig] = None,
    scheduler_config: Optional[SchedulerConfig] = None,
    cache_transceiver_config: Optional[CacheTransceiverConfig] = None,
    virtual_memory_pools: Optional[dict] = None,
    execution_stream: Optional[torch.cuda.Stream] = None,
    dwdp_manager: Optional[DwdpManager] = None,
    max_num_sequences: Optional[int] = None,
) -> PyExecutor:
    set_low_latency_dispatch(
        getattr(llm_args, 'enable_low_latency_host_dispatch', False))

    kv_cache_manager = resources.get(ResourceManagerType.KV_CACHE_MANAGER, None)

    spec_config = model_engine.spec_config

    is_disagg = is_disagg_enabled(cache_transceiver_config)

    max_num_sequences = resolve_max_num_sequences(
        model_engine,
        mapping,
        max_batch_size,
        llm_args,
        max_num_sequences=max_num_sequences)

    logger.info(
        f"max_seq_len={max_seq_len}, max_num_requests={max_num_sequences}, max_num_tokens={max_num_tokens}, max_batch_size={max_batch_size}"
    )
    for key, value in llm_args.extra_resource_managers.items():
        if key in resources:
            raise ValueError(
                f"Cannot overwrite existing resource manager {key}.")
        resources[key] = value

    peft_cache_manager = None
    if lora_config is not None:
        # TODO: Refactor dimension resolution into a LoraModuleDimensions
        # dataclass to avoid ad-hoc getattr + TP-division blocks per model type.
        from tensorrt_llm.bindings import LoraModule

        initial_lora_data_type = None
        if len(lora_config.lora_dir) == 1:
            # Route to appropriate loader based on checkpoint source
            initial_lora_data_type = _get_initial_lora_data_type(
                load_torch_lora(lora_config))
        else:
            assert len(lora_config.lora_target_modules
                       ) >= 1, "Expecting at least one lora target module"
            if not bool(lora_config.trtllm_modules_to_hf_modules):
                lora_config.trtllm_modules_to_hf_modules = get_default_trtllm_modules_to_hf_modules(
                )

        model_binding_config = model_engine.model.model_config.get_bindings_model_config(
            is_disagg=is_disagg)

        num_experts = _try_infer_num_experts(model_engine.model.model_config)

        num_kv_attention_heads_per_layer = model_binding_config.num_kv_heads_per_layer
        if max(num_kv_attention_heads_per_layer) != min(
                num_kv_attention_heads_per_layer):
            logger.warning(
                "Defining LORA with per-layer KV heads is not supported for LORA, using the max number of KV heads per layer"
            )
            num_kv_attention_heads = max(num_kv_attention_heads_per_layer)
        else:
            # all layers have the same number of KV heads
            num_kv_attention_heads = num_kv_attention_heads_per_layer[0]

        pretrained_config = model_engine.model.model_config.pretrained_config

        # Derive shared expert intermediate size from the LoRA adapter
        # weights, which are the source of truth for dimension validation.
        # The model config's shared_expert_intermediate_size may not match
        # the adapter (e.g., upcycled models).
        shared_expert_hidden_size = 0
        if lora_config.lora_dir:
            shared_expert_global = _infer_shared_expert_size_from_adapter(
                lora_config.lora_dir[0])
            if shared_expert_global > 0:
                shared_expert_hidden_size = shared_expert_global // mapping.tp_size

        moe_hidden_size = 0
        moe_intermediate = getattr(pretrained_config, 'moe_intermediate_size',
                                   None)
        if moe_intermediate is not None and moe_intermediate > 0:
            moe_hidden_size = moe_intermediate // mapping.tp_size

        # Mamba dimensions for hybrid models (e.g., Nemotron-H)
        # d_inner = mamba_head_dim * mamba_num_heads
        # d_in_proj = 2 * d_inner + 2 * n_groups * d_state + mamba_num_heads
        mamba_in_proj_size = 0
        mamba_inner_size = 0
        mamba_head_dim = getattr(pretrained_config, 'mamba_head_dim', 0)
        mamba_num_heads = getattr(pretrained_config, 'mamba_num_heads', 0)
        if mamba_head_dim > 0 and mamba_num_heads > 0:
            d_inner = mamba_head_dim * mamba_num_heads
            mamba_inner_size = d_inner // mapping.tp_size
            n_groups = getattr(pretrained_config, 'n_groups', 1)
            d_state = getattr(pretrained_config, 'ssm_state_size', 128)
            d_in_proj = 2 * d_inner + 2 * n_groups * d_state + mamba_num_heads
            mamba_in_proj_size = d_in_proj // mapping.tp_size

        # MoE latent size for latent MoE models (e.g., Nemotron-H SuperV3).
        # Latent projections are replicated (not TP-sharded), so pass the
        # raw config value without dividing by tp_size.
        moe_latent_size = getattr(pretrained_config, 'moe_latent_size', 0) or 0

        # For MoE models with shared experts: replace mlp_* target modules with
        # shared_expert_* equivalents. The shared expert uses different LoRA
        # module types with their own intermediate size. For pure MoE models
        # (no dense MLP layers), mlp_* modules don't exist in the model.
        target_modules = list(lora_config.lora_target_modules)
        if shared_expert_hidden_size > 0:
            has_dense_mlp = bool(
                getattr(pretrained_config, 'mlp_only_layers', None))
            mlp_to_shared_expert = {
                'mlp_h_to_4h': 'shared_expert_h_to_4h',
                'mlp_4h_to_h': 'shared_expert_4h_to_h',
                'mlp_gate': 'shared_expert_gate',
            }
            for mlp_mod, se_mod in mlp_to_shared_expert.items():
                if mlp_mod in target_modules:
                    if se_mod not in target_modules:
                        target_modules.append(se_mod)
                    if not has_dense_mlp:
                        target_modules.remove(mlp_mod)

        lora_modules = LoraModule.create_lora_modules(
            lora_module_names=target_modules,
            hidden_size=model_binding_config.hidden_size,
            mlp_hidden_size=model_binding_config.mlp_hidden_size,
            num_attention_heads=model_binding_config.num_heads,
            num_kv_attention_heads=num_kv_attention_heads,
            attention_head_size=model_binding_config.head_size,
            tp_size=mapping.tp_size,
            num_experts=num_experts,
            shared_expert_hidden_size=shared_expert_hidden_size,
            moe_hidden_size=moe_hidden_size,
            mamba_in_proj_size=mamba_in_proj_size,
            mamba_inner_size=mamba_inner_size,
            moe_latent_size=moe_latent_size)

        model_binding_config.use_lora_plugin = True
        model_binding_config.lora_modules = lora_modules
        model_binding_config.max_lora_rank = lora_config.max_lora_rank

        max_lora_rank = lora_config.max_lora_rank
        num_lora_modules = _compute_num_lora_modules(
            pretrained_config,
            target_modules + lora_config.missing_qkv_modules,
        )

        peft_cache_config_model = PeftCacheConfig(
        ) if peft_cache_config is None else peft_cache_config
        if lora_config.max_loras is not None:
            peft_cache_config_model.num_device_module_layer = \
                max_lora_rank * num_lora_modules * lora_config.max_loras
        if lora_config.max_cpu_loras is not None:
            peft_cache_config_model.num_host_module_layer = \
                max_lora_rank * num_lora_modules * lora_config.max_cpu_loras

        from tensorrt_llm.bindings import WorldConfig
        world_config = WorldConfig(
            tensor_parallelism=mapping.tp_size,
            pipeline_parallelism=mapping.pp_size,
            context_parallelism=mapping.cp_size,
            rank=dist.mapping.rank,
            gpus_per_node=dist.mapping.gpus_per_node,
        )
        peft_cache_manager = PeftCacheManager(
            peft_cache_config=peft_cache_config_model,
            lora_config=lora_config,
            model_config=model_binding_config,
            world_config=world_config,
            execution_stream=execution_stream,
            lora_target_modules=target_modules,
            initial_data_type=initial_lora_data_type,
        )
        resources[ResourceManagerType.PEFT_CACHE_MANAGER] = peft_cache_manager
        model_engine.set_lora_model_config(
            target_modules, lora_config.trtllm_modules_to_hf_modules,
            lora_config.swap_gate_up_proj_lora_b_weight)
        if isinstance(model_engine, PyTorchModelEngine):
            model_engine._init_cuda_graph_lora_manager(lora_config)

    validate_seq_slot_pool_covers_admission(max_num_sequences, kv_cache_manager)
    resources[ResourceManagerType.SEQ_SLOT_MANAGER] = SeqSlotManager(
        max_num_sequences)

    compression_manager = resources.get(
        ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER)
    if compression_manager is not None:
        compression_manager.bind_kv_cache_managers(
            resources[ResourceManagerType.KV_CACHE_MANAGER],
            resources.get(ResourceManagerType.DRAFT_KV_CACHE_MANAGER),
        )
    resource_manager = ResourceManager(resources)

    # KV cache manager runs last (others may depend on it), except the
    # compression manager which reconciles after it (below).
    if kv_cache_manager is not None:
        resource_manager.resource_managers.move_to_end(
            ResourceManagerType.KV_CACHE_MANAGER, last=True)
    cross_kv_cache_manager = resources.get(
        ResourceManagerType.CROSS_KV_CACHE_MANAGER)
    if cross_kv_cache_manager is not None:
        resource_manager.resource_managers.move_to_end(
            ResourceManagerType.CROSS_KV_CACHE_MANAGER, last=True)
    # Iteration-driven compression is the final reconciler after every native
    # KV manager. Cold-page quantization runs only at native storage migration.
    if (ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER
            in resource_manager.resource_managers):
        resource_manager.resource_managers.move_to_end(
            ResourceManagerType.KV_CACHE_COMPRESSION_MANAGER, last=True)

    # When scheduler_capacity == 1, attention dp dummy request will prevent the scheduling of DISAGG_GENERATION_INIT.
    # Enlarge scheduler capacity to avoid DISAGG_GENERATION_INIT stuck in the scheduler.
    # V1 scheduler handles overlap via two_step_lookahead, so skip the
    # slot-pool overlap factor here.
    scheduler_capacity = max_batch_size * mapping.pp_size
    if scheduler_capacity == 1 and mapping.enable_attention_dp and kv_cache_manager:
        scheduler_capacity += 1

    # For encoder-decoder models, requests start in ENCODER_INIT and the
    # capacity scheduler must admit them already at that state so the
    # encoder loop can run. Decoder-only deployments keep the default
    # CONTEXT_INIT gating.
    no_schedule_until_state = (LlmRequestState.ENCODER_INIT
                               if cross_kv_cache_manager is not None else
                               LlmRequestState.CONTEXT_INIT)

    # V2 scheduler uses scheduler_capacity as the per-iteration request
    # budget (BudgetTracker.max_num_requests).  Unlike V1 which has a
    # separate CapacityScheduler (needs pp_size * max_batch_size to hold
    # requests across PP stages) and MicroBatchScheduler (uses
    # max_batch_size for per-forward batch limit), V2 merges both into
    # one loop.  PP on-the-fly is handled by inflight_request_ids
    # filtering, so its budget should be based on max_batch_size, not
    # max_num_sequences (which includes the pp_size multiplier).
    v2_scheduler_capacity = max_batch_size
    if v2_scheduler_capacity == 1 and mapping.enable_attention_dp and kv_cache_manager:
        v2_scheduler_capacity += 1

    if isinstance(kv_cache_manager, KVCacheManagerV2):
        # V2: interleaved scheduler handles both capacity and budget
        draft_kv_cache_manager = resources.get(
            ResourceManagerType.DRAFT_KV_CACHE_MANAGER)
        scheduler_policy = (scheduler_config.capacity_scheduler_policy
                            if scheduler_config is not None else
                            CapacitySchedulerPolicy.MAX_UTILIZATION)
        enable_prefix_aware_scheduling = (
            scheduler_config.enable_prefix_aware_scheduling
            if scheduler_config is not None else True)
        scheduler = KVCacheV2Scheduler(
            max_batch_size=max_batch_size,
            max_num_tokens=max_num_tokens,
            kv_cache_manager=kv_cache_manager,
            scheduler_policy=scheduler_policy,
            ctx_chunk_config=ctx_chunk_config,
            peft_cache_manager=peft_cache_manager.impl
            if peft_cache_manager is not None else None,
            scheduler_capacity=v2_scheduler_capacity,
            draft_kv_cache_manager=draft_kv_cache_manager,
            cross_kv_cache_manager=cross_kv_cache_manager,
            no_schedule_until_state=no_schedule_until_state,
            enable_prefix_aware_scheduling=enable_prefix_aware_scheduling,
            # A disaggregated generation worker must not replay context locally.
            enable_recompute_pause=not is_disagg,
        )
    elif (scheduler_config is not None
          and scheduler_config.use_python_scheduler):
        enable_prefix_aware_scheduling = scheduler_config.enable_prefix_aware_scheduling
        scheduler = SimpleUnifiedScheduler(
            max_batch_size=max_batch_size,
            max_num_tokens=max_num_tokens,
            kv_cache_manager=kv_cache_manager.impl
            if kv_cache_manager is not None else None,
            peft_cache_manager=peft_cache_manager.impl
            if peft_cache_manager is not None else None,
            scheduler_policy=scheduler_config.capacity_scheduler_policy,
            ctx_chunk_config=ctx_chunk_config,
            cross_kv_cache_manager=cross_kv_cache_manager.impl
            if cross_kv_cache_manager is not None else None,
            two_step_lookahead=mapping.has_pp(),
            scheduler_capacity=scheduler_capacity,
            no_schedule_until_state=no_schedule_until_state,
            enable_prefix_aware_scheduling=enable_prefix_aware_scheduling,
        )
    else:
        enable_prefix_aware_scheduling = (
            scheduler_config.enable_prefix_aware_scheduling
            if scheduler_config is not None else True)
        capacity_scheduler = BindCapacityScheduler(
            scheduler_capacity,
            kv_cache_manager.impl if kv_cache_manager is not None else None,
            peft_cache_manager.impl if peft_cache_manager is not None else None,
            scheduler_config.capacity_scheduler_policy,
            cross_kv_cache_manager=cross_kv_cache_manager.impl
            if cross_kv_cache_manager is not None else None,
            two_step_lookahead=mapping.has_pp(),
            no_schedule_until_state=no_schedule_until_state,
            enable_prefix_aware_scheduling=enable_prefix_aware_scheduling,
        )

        mb_scheduler = BindMicroBatchScheduler(
            max_batch_size,
            max_num_tokens,
            ctx_chunk_config,
            no_schedule_until_state=no_schedule_until_state,
        )

        reorder_policy_config = llm_args.reorder_policy_config
        if reorder_policy_config is not None:
            assert reorder_policy_config.policy_name == "AgentTree", "Reorder policy only supports AgentTree for now"
            capacity_scheduler.impl.set_agent_tree_reorder_policy(
                reorder_policy_config.policy_args.agent_percentage,
                reorder_policy_config.policy_args.agent_types,
                reorder_policy_config.policy_args.agent_inflight_seq_num)
        scheduler = SimpleScheduler(capacity_scheduler, mb_scheduler)

    if getattr(model_engine, "mm_encoder_item_scheduling_enabled", False):
        # Wrap the LLM scheduler with atomic MM item budgeting. ModelEngine
        # already validated model-capability-dependent feature combinations.
        multimodal_config = llm_args.multimodal_config
        # `mm_encoder_item_scheduling_enabled` already excludes the DISABLED
        # policy (a disabled model never reaches here and keeps the base LLM
        # scheduler), so only the EAGER vs DEFAULT variant is selected here.
        scheduler_cls = MultimodalScheduler
        if (multimodal_config.encoder_scheduling_policy ==
                MultimodalEncoderSchedulingPolicy.EAGER):
            logger.info("Eager multimodal encoder scheduling is enabled for "
                        "capacity-rejected active requests.")
            scheduler_cls = MultimodalEagerEncoderScheduler
        scheduler = scheduler_cls(
            scheduler,
            max_batch_size=model_engine.encoder_batch_size,
            max_num_tokens=model_engine.encoder_max_num_tokens,
            output_budget_bytes=model_engine.mm_encoder_output_budget_bytes,
            bytes_per_encoder_embedding=(
                model_engine.bytes_per_mm_encoder_embedding),
        )

    config = model_engine.model.model_config.pretrained_config
    attention_type = AttentionTypeCpp.MLA if is_mla(
        config) else AttentionTypeCpp.DEFAULT

    # For hybrid models, this has both impl and mamba_impl
    mamba_cache_manager = None
    if isinstance(kv_cache_manager, BaseMambaCacheManager):
        mamba_cache_manager = kv_cache_manager

    kv_cache_transceiver = create_kv_cache_transceiver(
        mapping, dist, kv_cache_manager, attention_type,
        cache_transceiver_config, mamba_cache_manager)

    waiting_queue_policy = (scheduler_config.waiting_queue_policy
                            if scheduler_config is not None else
                            WaitingQueuePolicy.FCFS)

    # For enc-dec models max_seq_len covers the (longer) encoder sequence, so
    # cap the executor's per-request max_tokens at the decoder position table
    # (max_target_positions).
    executor_max_seq_len = max_seq_len
    if model_engine.model.model_config.is_encoder_decoder:
        decoder_position_limit = getattr(config, "max_target_positions", None)
        if (decoder_position_limit is not None
                and executor_max_seq_len is not None):
            executor_max_seq_len = min(executor_max_seq_len,
                                       int(decoder_position_limit))

    return PyExecutor(
        resource_manager,
        scheduler,
        model_engine=model_engine,
        sampler=sampler,
        drafter=drafter,
        dist=dist,
        max_num_sequences=max_num_sequences,
        disable_overlap_scheduler=llm_args.disable_overlap_scheduler,
        enable_early_first_token_response=llm_args.
        enable_early_first_token_response,
        max_batch_size=max_batch_size,
        max_beam_width=max_beam_width,
        max_draft_len=spec_config.max_draft_len
        if spec_config is not None else 0,
        max_total_draft_tokens=(spec_config.tokens_per_gen_step -
                                1) if spec_config is not None else 0,
        kv_cache_transceiver=kv_cache_transceiver,
        guided_decoder=guided_decoder,
        start_worker=start_worker,
        garbage_collection_gen0_threshold=garbage_collection_gen0_threshold,
        kv_connector_manager=kv_connector_manager,
        resource_governor_queue=resource_governor_queue,
        max_seq_len=executor_max_seq_len,
        peft_cache_config=peft_cache_config,
        virtual_memory_pools=virtual_memory_pools,
        execution_stream=execution_stream,
        waiting_queue_policy=waiting_queue_policy,
        dwdp_manager=dwdp_manager,
        enable_kv_pool_rebalance=llm_args.kv_cache_config.
        enable_kv_pool_rebalance,
    )


def create_torch_sampler_args(
    *,
    max_seq_len: int,
    speculative_config: SpeculativeConfig,
    max_beam_width: int,
    disable_overlap_scheduler: bool,
    enable_async_worker: bool,
    enable_speculative_beam_history_d2h: bool,
    max_num_sequences: int,
):
    # The sampler's per-slot state is indexed by sequence slots, so it must
    # be sized identically to the executor's slot pool.
    max_draft_len = (0 if speculative_config is None else
                     speculative_config.max_draft_len)
    max_total_draft_tokens = (0 if speculative_config is None else
                              speculative_config.tokens_per_gen_step - 1)

    return TorchSampler.Args(
        max_seq_len=max_seq_len,
        max_draft_len=max_draft_len,
        max_total_draft_tokens=max_total_draft_tokens,
        max_num_sequences=max_num_sequences,
        max_beam_width=max_beam_width,
        disable_overlap_scheduler=disable_overlap_scheduler,
        enable_async_worker=enable_async_worker,
        enable_speculative_beam_history_d2h=enable_speculative_beam_history_d2h,
    )


def instantiate_sampler(
    engine: PyTorchModelEngine,
    llm_args: TorchLlmArgs,
    mapping: Mapping,
    *,
    max_batch_size: int,
    max_beam_width: int,
    mm_encoder_only: bool,
    speculative_config: SpeculativeConfig,
    max_num_sequences: Optional[int] = None,
):
    enable_async_worker = (confidential_compute_enabled()
                           or llm_args.sampler_force_async_worker)

    max_num_sequences = resolve_max_num_sequences(
        engine,
        mapping,
        max_batch_size,
        llm_args,
        max_num_sequences=max_num_sequences)

    sampler_args = create_torch_sampler_args(
        max_seq_len=engine.max_seq_len,
        speculative_config=speculative_config,
        max_beam_width=max_beam_width,
        disable_overlap_scheduler=llm_args.disable_overlap_scheduler,
        enable_async_worker=enable_async_worker,
        enable_speculative_beam_history_d2h=llm_args.
        enable_speculative_beam_history_d2h,
        max_num_sequences=max_num_sequences,
    )
    if engine.spec_config is not None and engine.spec_config.spec_dec_mode.has_spec_decoder(
    ):
        return get_spec_decoder(sampler_args, engine.spec_config)

    if mm_encoder_only:
        # NOTE: handle model outputs specially for mm encoder executor/engine
        return EarlyStopWithMMResult()
    if not engine.model.model_config.is_generation:
        # NOTE: choose sampler based on model type
        return EarlyStopSampler()
    return TorchSampler(sampler_args)


_ATTN_MODULES = frozenset({
    "attn_q",
    "attn_k",
    "attn_v",
    "attn_qkv",
    "attn_dense",
    "cross_attn_q",
    "cross_attn_k",
    "cross_attn_v",
})
_MLP_MODULES = frozenset({
    "mlp_h_to_4h",
    "mlp_4h_to_h",
    "mlp_gate",
    "mlp_gate_up",
})


def _compute_num_lora_modules(pretrained_config,
                              all_target_modules: list[str]) -> int:
    """Compute the total number of LoRA module-layer slots for cache sizing.

    For models with per-layer block_configs (e.g. Nemotron-NAS / DeciLM),
    layers with no_op or replace_with_linear attention/FFN cannot host LoRA
    adapters, so they are excluded from the count.  For all other models,
    falls back to the uniform num_hidden_layers x len(target_modules).
    """
    num_layers = pretrained_config.num_hidden_layers
    block_configs = getattr(pretrained_config, "block_configs", None)

    if block_configs is None:
        return num_layers * len(all_target_modules)

    attn_modules = [m for m in all_target_modules if m in _ATTN_MODULES]
    mlp_modules = [m for m in all_target_modules if m in _MLP_MODULES]
    other_modules = [
        m for m in all_target_modules
        if m not in _ATTN_MODULES and m not in _MLP_MODULES
    ]

    def _has_lora_capable_attn(bc):
        return not bc.attention.no_op and not bc.attention.replace_with_linear

    def _has_lora_capable_ffn(bc):
        return not bc.ffn.no_op and not bc.ffn.replace_with_linear

    layers_with_attn = sum(1 for bc in block_configs
                           if _has_lora_capable_attn(bc))
    layers_with_mlp = sum(1 for bc in block_configs
                          if _has_lora_capable_ffn(bc))

    total = (layers_with_attn * len(attn_modules) +
             layers_with_mlp * len(mlp_modules) +
             num_layers * len(other_modules))

    logger.info(f"LoRA module-layer count: {total} "
                f"(attn: {layers_with_attn}x{len(attn_modules)}, "
                f"mlp: {layers_with_mlp}x{len(mlp_modules)}, "
                f"other: {num_layers}x{len(other_modules)}, "
                f"uniform would be {num_layers * len(all_target_modules)})")
    return total


def _infer_shared_expert_size_from_adapter(adapter_dir: str) -> int:
    """Infer shared expert intermediate size from LoRA adapter weights.

    Scans the adapter for shared_expert.down_proj lora_A weights and
    returns the global (unsharded) intermediate size. This is more reliable
    than the model config, which may not match the adapter (e.g., upcycled
    models).
    """
    import json

    try:
        from tensorrt_llm.models.convert_utils import (get_model_path,
                                                       load_state_dict)
        model_path = get_model_path(adapter_dir, "adapter_model")
        if model_path is None:
            return 0
        adapter_weights = load_state_dict(model_path)
        if adapter_weights is None:
            return 0
        for key, tensor in adapter_weights.items():
            if 'shared_expert' in key and 'down_proj' in key and 'lora_A' in key:
                adapter_config_path = os.path.join(adapter_dir,
                                                   "adapter_config.json")
                with open(adapter_config_path) as f:
                    rank = json.load(f).get("r", 0)
                if rank > 0:
                    return tensor.shape[1] if tensor.shape[
                        0] == rank else tensor.shape[0]
    except Exception as e:
        logger.debug(f"Failed to infer shared expert size from adapter: {e}")
    return 0


def _try_infer_num_experts(model_config: ModelConfig) -> int:
    """
    Attempt to infer the number of experts from the model configuration.

    Different MoE models use different attribute names for storing the number of experts,
    so this function checks for various possible names and returns a default of 1 if none are found.
    However, this function is not exhaustive and may miss some cases, so it should be revised.
    """
    config = getattr(model_config, 'pretrained_config', model_config)

    expert_attr_names = [
        'num_experts', 'num_local_experts', 'moe_num_experts',
        'experts_per_router'
    ]
    num_experts = None
    for attr_name in expert_attr_names:
        if hasattr(config, attr_name):
            num_experts = getattr(config, attr_name)
            break

    # Default to 1 for non-MoE models or if no experts attribute is found
    if num_experts is None:
        return 1

    return num_experts


def _adjust_torch_mem_fraction():
    # If true, adjust PyTorch CUDA memory fraction to correspond to the
    # total GPU memory minus the statically allocated engine memory.
    # If false, set the PyTorch CUDA memory fraction to 1.0.
    _limit_torch_cuda_mem_fraction: bool = True

    # FIXME: PyTorch only uses the garbage_collection_threshold setting
    #        if a memory fraction is set, cf.
    #   https://github.com/pytorch/pytorch/blob/cd995bfb2aac8891465809be3ce29543bd524287/c10/cuda/CUDACachingAllocator.cpp#L1357
    logger.debug("Setting PyTorch memory fraction to 1.0")
    torch.cuda.set_per_process_memory_fraction(1.0)

    # FIXME: As soon as
    #     torch.cuda._set_allocator_settings (added in PyTorch 2.8.0-rc1)
    #   or a similar API is available, the warning below should be removed
    #   and the allocator GC threshold be set via the new API instead.
    torch_allocator_config = os.environ.get("PYTORCH_ALLOC_CONF", "")
    torch_mem_threshold_advised = (
        torch.cuda.get_allocator_backend() == "native"
        and "expandable_segments:True" not in torch_allocator_config)
    torch_mem_threshold_set = "garbage_collection_threshold:" in torch_allocator_config
    if torch_mem_threshold_advised and not torch_mem_threshold_set:
        logger.warning(
            "It is recommended to incl. 'garbage_collection_threshold:0.???' or 'backend:cudaMallocAsync'"
            " or 'expandable_segments:True' in PYTORCH_ALLOC_CONF.")

    # NOTE: Even if a memory threshold was not set (cf. warning above), setting a memory
    #       fraction < 1.0 is beneficial, because
    #         https://github.com/pytorch/pytorch/blob/5228986c395dc79f90d2a2b991deea1eef188260/c10/cuda/CUDACachingAllocator.cpp#L2719
    #       and
    #         https://github.com/pytorch/pytorch/blob/5228986c395dc79f90d2a2b991deea1eef188260/c10/cuda/CUDACachingAllocator.cpp#L1240
    #       lead PyTorch to release all unused memory before hitting the set fraction. This
    #       still mitigates OOM, although at a higher performance impact, because it
    #       effectively resets the allocator cache.
    if not _limit_torch_cuda_mem_fraction:
        return
    mem_reserved = torch.cuda.memory_reserved()
    mem_free, mem_total = torch.cuda.mem_get_info()
    safety_margin = 32 * 1024**2
    mem_torch_max = mem_free + mem_reserved - safety_margin
    mem_torch_fraction = mem_torch_max / mem_total
    logger.info(
        f"Setting PyTorch memory fraction to {mem_torch_fraction} ({mem_torch_max / 1024**3} GiB)"
    )
    torch.cuda.set_per_process_memory_fraction(mem_torch_fraction)


def validate_feature_combination(llm_args, model_engine):
    # Validate the flags for features' combination
    def init_feature_status(llm_args) -> Dict[str, bool]:
        assert isinstance(
            llm_args, TorchLlmArgs
        ), "Expect TorchLlmArgs used for feature status validation."
        feature_list = [
            "overlap_scheduler",
            "cuda_graph",
            "attention_dp",
            "disaggregated_serving",
            "chunked_prefill",
            "mtp",
            "eagle3_one_model",
            "kv_cache_reuse",
            "slide_window_attention",
            "guided_decoding",
        ]
        feature_status: Dict[str, bool] = dict.fromkeys(feature_list)
        feature_status[
            "overlap_scheduler"] = not llm_args.disable_overlap_scheduler
        feature_status["cuda_graph"] = llm_args.cuda_graph_config is not None
        feature_status["attention_dp"] = llm_args.enable_attention_dp
        feature_status[
            "disaggregated_serving"] = llm_args.cache_transceiver_config is not None
        feature_status["chunked_prefill"] = llm_args.enable_chunked_prefill
        feature_status["mtp"] = isinstance(llm_args.speculative_config,
                                           MTPDecodingConfig)
        feature_status["eagle3_one_model"] = isinstance(
            llm_args.speculative_config, EagleDecodingConfig)
        feature_status[
            "kv_cache_reuse"] = llm_args.kv_cache_config is not None and llm_args.kv_cache_config.enable_block_reuse
        feature_status["slide_window_attention"] = (
            hasattr(model_engine.model.model_config.pretrained_config,
                    "layer_types") and "sliding_attention"
            in model_engine.model.model_config.pretrained_config.layer_types)
        feature_status[
            "guided_decoding"] = llm_args.guided_decoding_backend is not None
        assert all(v is not None for v in feature_status.values()
                   ), "feature status has not been fully initialized."
        assert all(
            k in feature_list for k in
            feature_status.keys()), "unexpected feature type in feature_status."
        return feature_status

    feature_status: Dict[str, bool] = init_feature_status(llm_args)

    # Kept as an extension point; there are currently no conflicting feature
    # combinations to reject.
    CONFLICT_RULES: list[dict[str, Any]] = [
        # Add new conflict rules here in the future
    ]
    for rule in CONFLICT_RULES:
        if all(feature_status[feature] for feature in rule["features"]):
            raise ValueError(rule["message"])
