File size: 158,191 Bytes
6d93aeb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
"""Generated RWKV-7 inference runtime. Do not edit; regenerate it."""

import linecache
import sys
from types import ModuleType

RUNTIME_FORMAT_VERSION = 4
SOURCE_SHA256 = {'configuration_rwkv7': '60a060199526f5a6c4d5a2436520eed97e9f21f2ea3cb6e6584d9b5fcb6f21f4', 'custom_ops': 'fb86e5c9dfc9deb5a1be821ce02c3a390ce5c95ff55feded6add2422f0dc3623', 'kernel_dispatch': '0290559437016ba4e2c2a26f7c7a0985b6fbab991535b7ae65ba18eb421fa9e1', 'modeling_rwkv7': '7e1e57a4516efe78a74f971ada0fe58cb71c87102479ed7516b6d3585964ca21', 'state': 'd13d81af6a301a70024494f03b95006bf5f2471ba0b57fd7f1f94bd54f70f2ef', 'tilelang_decode': 'fe63e6332037cacbd0400d82f65fd6f1c282e31d92906c5f34d39eb9e83eb642'}
_SOURCES = {'configuration_rwkv7': 'from __future__ import annotations\n\nfrom typing import Any, Mapping\n\nimport torch\nfrom transformers import PretrainedConfig\n\n\nclass RWKV7Config(PretrainedConfig):\n    """Configuration for RWKV-7 causal language models."""\n\n    model_type = "rwkv7"\n    keys_to_ignore_at_inference = ("past_key_values",)\n\n    def __init__(\n        self,\n        vocab_size: int = 65536,\n        hidden_size: int = 768,\n        num_hidden_layers: int = 12,\n        head_size: int = 64,\n        intermediate_size: int | None = None,\n        decay_lora_rank: int = 64,\n        a_lora_rank: int = 64,\n        gate_lora_rank: int = 128,\n        value_lora_rank: int = 32,\n        layer_norm_epsilon: float = 1e-5,\n        use_cache: bool = True,\n        kernel_backend: str = "auto",\n        recurrent_state_dtype: str = "float32",\n        rwkv_prefill_chunk_size: int = 256,\n        initializer_range: float = 0.02,\n        tie_word_embeddings: bool = False,\n        bos_token_id: int | None = None,\n        eos_token_id: int | None = 0,\n        pad_token_id: int | None = 0,\n        **kwargs: Any,\n    ) -> None:\n        if hidden_size % head_size:\n            raise ValueError("hidden_size must be divisible by head_size")\n        if kernel_backend not in {"auto", "torch", "tilelang"}:\n            raise ValueError(f"Unsupported kernel backend: {kernel_backend}")\n        if recurrent_state_dtype not in {"float32", "float16", "bfloat16"}:\n            raise ValueError(\n                "recurrent_state_dtype must be float32, float16, or bfloat16"\n            )\n        if rwkv_prefill_chunk_size < 0:\n            raise ValueError("rwkv_prefill_chunk_size must be non-negative")\n        self.vocab_size = vocab_size\n        self.hidden_size = hidden_size\n        self.num_hidden_layers = num_hidden_layers\n        self.head_size = head_size\n        self.num_attention_heads = hidden_size // head_size\n        self.intermediate_size = intermediate_size or hidden_size * 4\n        self.decay_lora_rank = decay_lora_rank\n        self.a_lora_rank = a_lora_rank\n        self.gate_lora_rank = gate_lora_rank\n        self.value_lora_rank = value_lora_rank\n        self.layer_norm_epsilon = layer_norm_epsilon\n        self.use_cache = use_cache\n        self.kernel_backend = kernel_backend\n        self.recurrent_state_dtype = recurrent_state_dtype\n        self.rwkv_prefill_chunk_size = rwkv_prefill_chunk_size\n        self.initializer_range = initializer_range\n        self.is_decoder = True\n        self.is_encoder_decoder = False\n        super().__init__(\n            tie_word_embeddings=tie_word_embeddings,\n            bos_token_id=bos_token_id,\n            eos_token_id=eos_token_id,\n            pad_token_id=pad_token_id,\n            **kwargs,\n        )\n\n    @classmethod\n    def from_rwkv_weights(\n        cls, weights: Mapping[str, torch.Tensor], **kwargs: Any\n    ) -> "RWKV7Config":\n        embedding = weights["emb.weight"]\n        hidden_size = embedding.shape[1]\n        layer_ids: set[int] = set()\n        for name in weights:\n            if not name.startswith("blocks."):\n                continue\n            try:\n                layer_ids.add(int(name.split(".")[1]))\n            except (IndexError, ValueError):\n                continue\n        if not layer_ids:\n            raise ValueError("Checkpoint contains no RWKV blocks")\n        head_size = weights["blocks.0.att.r_k"].shape[1]\n        inferred: dict[str, Any] = {\n            "vocab_size": embedding.shape[0],\n            "hidden_size": hidden_size,\n            "num_hidden_layers": max(layer_ids) + 1,\n            "head_size": head_size,\n            "intermediate_size": weights["blocks.0.ffn.key.weight"].shape[0],\n            "decay_lora_rank": weights["blocks.0.att.w1"].shape[1],\n            "a_lora_rank": weights["blocks.0.att.a1"].shape[1],\n            "gate_lora_rank": weights["blocks.0.att.g1"].shape[1],\n            "value_lora_rank": weights["blocks.1.att.v1"].shape[1]\n            if max(layer_ids) >= 1\n            else max(1, hidden_size // 24),\n        }\n        inferred.update(kwargs)\n        return cls(**inferred)\n', 'custom_ops': 'from __future__ import annotations\n\nfrom typing import Any\n\nimport torch\n\nFLOAT32 = torch.float32  # type: ignore[attr-defined]\n\n\ndef _validate_state_finalize_inputs(\n    state: torch.Tensor,\n    decay: torch.Tensor,\n    anti_update: torch.Tensor,\n    value_key: torch.Tensor,\n) -> None:\n    if state.ndim != 4:\n        raise ValueError("state must have shape [batch, heads, head_size, head_size]")\n    batch_size, num_heads, rows, columns = state.shape\n    if rows != columns:\n        raise ValueError("state matrices must be square")\n    if decay.shape != (batch_size, num_heads, columns):\n        raise ValueError("decay must have shape [batch, heads, head_size]")\n    if anti_update.shape != state.shape:\n        raise ValueError("anti_update must match state shape")\n    if value_key.shape != state.shape:\n        raise ValueError("value_key must match state shape")\n    tensors = (state, decay, anti_update, value_key)\n    if any(tensor.device.type != "cuda" for tensor in tensors):\n        raise ValueError("state finalize custom op requires CUDA tensors")\n    if any(tensor.device != state.device for tensor in tensors[1:]):\n        raise ValueError("state finalize tensors must use one CUDA device")\n    if state.dtype != FLOAT32:\n        raise ValueError("state must use float32")\n    if anti_update.dtype != FLOAT32 or value_key.dtype != FLOAT32:\n        raise ValueError("anti_update and value_key must use float32")\n    supported_decay_dtypes = {\n        FLOAT32,\n        getattr(torch, "float16"),\n        getattr(torch, "bfloat16"),\n    }\n    if decay.dtype not in supported_decay_dtypes:\n        raise ValueError("decay must use float32, float16, or bfloat16")\n\n\ndef _validate_state_finalize_backward_inputs(\n    grad_output: torch.Tensor, decay: torch.Tensor\n) -> None:\n    if grad_output.ndim != 4 or grad_output.dtype != FLOAT32:\n        raise ValueError("grad_output must be a rank-4 float32 tensor")\n    batch_size, num_heads, rows, columns = grad_output.shape\n    if rows != columns or decay.shape != (batch_size, num_heads, columns):\n        raise ValueError("state finalize backward shapes are incompatible")\n    if grad_output.device != decay.device:\n        raise ValueError("state finalize backward tensors must use one CUDA device")\n    if not grad_output.is_cuda or not decay.is_cuda:\n        raise ValueError("state finalize backward custom op requires CUDA tensors")\n\n\ndef _register_state_finalize_backward_op() -> Any:\n    namespace = torch.ops.rwkv7_pytorch\n    if hasattr(namespace, "_state_finalize_backward"):\n        return namespace._state_finalize_backward.default\n\n    @torch.library.custom_op(\n        "rwkv7_pytorch::_state_finalize_backward",\n        mutates_args=(),\n        device_types="cuda",\n    )\n    def state_finalize_backward(\n        grad_output: torch.Tensor,\n        decay: torch.Tensor,\n    ) -> torch.Tensor:\n        _validate_state_finalize_backward_inputs(grad_output, decay)\n        from .kernel import state as _kernel_state\n        tilelang_state_finalize_backward = _kernel_state.tilelang_state_finalize_backward\n\n        return tilelang_state_finalize_backward(grad_output, decay)\n\n    @state_finalize_backward.register_fake\n    def state_finalize_backward_fake(\n        grad_output: torch.Tensor, decay: torch.Tensor\n    ) -> torch.Tensor:\n        del decay\n        return grad_output.new_empty(grad_output.shape)\n\n    return state_finalize_backward\n\n\n_STATE_FINALIZE_BACKWARD_OP = _register_state_finalize_backward_op()\n\n\ndef _register_state_finalize_op() -> Any:\n    namespace = torch.ops.rwkv7_pytorch\n    if hasattr(namespace, "state_finalize"):\n        return namespace.state_finalize.default\n\n    @torch.library.custom_op(\n        "rwkv7_pytorch::state_finalize",\n        mutates_args=(),\n        device_types="cuda",\n    )\n    def state_finalize(\n        state: torch.Tensor,\n        decay: torch.Tensor,\n        anti_update: torch.Tensor,\n        value_key: torch.Tensor,\n    ) -> torch.Tensor:\n        _validate_state_finalize_inputs(state, decay, anti_update, value_key)\n        from .kernel import state as _kernel_state\n        tilelang_state_finalize = _kernel_state.tilelang_state_finalize\n\n        return tilelang_state_finalize(state, decay, anti_update, value_key)\n\n    @state_finalize.register_fake\n    def state_finalize_fake(\n        state: torch.Tensor,\n        decay: torch.Tensor,\n        anti_update: torch.Tensor,\n        value_key: torch.Tensor,\n    ) -> torch.Tensor:\n        del decay, anti_update, value_key\n        return state.new_empty(state.shape)\n\n    def setup_context(ctx: Any, inputs: tuple[torch.Tensor, ...], output: torch.Tensor) -> None:\n        del output\n        state, decay, _, _ = inputs\n        ctx.save_for_backward(state, decay)\n\n    def backward(\n        ctx: Any, grad_output: torch.Tensor\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n        state, decay = ctx.saved_tensors\n        if torch.is_grad_enabled():\n            grad_state = grad_output * decay.float().unsqueeze(-2)\n        else:\n            grad_state = _STATE_FINALIZE_BACKWARD_OP(grad_output, decay)\n        grad_decay = (grad_output * state).sum(dim=-2).to(decay.dtype)\n        return grad_state, grad_decay, grad_output, grad_output\n\n    state_finalize.register_autograd(backward, setup_context=setup_context)\n    return state_finalize\n\n\n_STATE_FINALIZE_OP = _register_state_finalize_op()\n\n\ndef tilelang_state_finalize_op(\n    state: torch.Tensor,\n    decay: torch.Tensor,\n    anti_update: torch.Tensor,\n    value_key: torch.Tensor,\n) -> torch.Tensor:\n    """Run differentiable TileLang state finalization through PyTorch dispatcher."""\n    _validate_state_finalize_inputs(state, decay, anti_update, value_key)\n    return _STATE_FINALIZE_OP(state, decay, anti_update, value_key)\n', 'kernel_dispatch': 'from __future__ import annotations\n\nfrom contextlib import suppress\nfrom dataclasses import dataclass\nfrom functools import lru_cache\nfrom importlib.metadata import PackageNotFoundError\nfrom importlib.metadata import version as distribution_version\nfrom importlib.util import find_spec\nfrom typing import Any\n\nimport torch\n\nSUPPORTED_TILELANG_VERSIONS = ("0.1.12",)\n\n\n@dataclass(frozen=True)\nclass KernelBackendStatus:\n    name: str\n    available: bool\n    version: str | None = None\n    reason: str | None = None\n\n\n@lru_cache(maxsize=1)\ndef tilelang_status() -> KernelBackendStatus:\n    if find_spec("tilelang") is None:\n        return KernelBackendStatus("tilelang", False, reason="package not installed")\n    installed_version = None\n    with suppress(PackageNotFoundError):\n        installed_version = distribution_version("tilelang")\n    if installed_version is None:\n        return KernelBackendStatus(\n            "tilelang", False, reason="distribution metadata unavailable"\n        )\n    if installed_version not in SUPPORTED_TILELANG_VERSIONS:\n        supported = ", ".join(SUPPORTED_TILELANG_VERSIONS)\n        return KernelBackendStatus(\n            "tilelang",\n            False,\n            version=installed_version,\n            reason=f"unsupported version; expected one of: {supported}",\n        )\n    if not torch.cuda.is_available():\n        return KernelBackendStatus(\n            "tilelang", False, version=installed_version, reason="CUDA unavailable"\n        )\n    return KernelBackendStatus("tilelang", True, version=installed_version)\n\n\ndef kernel_backend_status() -> dict[str, KernelBackendStatus]:\n    return {\n        "torch": KernelBackendStatus("torch", True, version=torch.__version__),\n        "tilelang": tilelang_status(),\n    }\n\n\ndef available_backends() -> tuple[str, ...]:\n    return tuple(\n        name for name, status in kernel_backend_status().items() if status.available\n    )\n\n\ndef resolve_backend(requested: str, device: Any) -> str:\n    if requested not in {"auto", "torch", "tilelang"}:\n        raise ValueError(f"Unsupported kernel backend: {requested}")\n    normalized = requested\n    if normalized == "auto":\n        device_type = getattr(device, "type", str(device).split(":", 1)[0])\n        if device_type == "cuda" and tilelang_status().available:\n            return "tilelang"\n        return "torch"\n    status = kernel_backend_status()[normalized]\n    if not status.available:\n        detail = f": {status.reason}" if status.reason else ""\n        raise RuntimeError(f"Kernel backend is unavailable: {normalized}{detail}")\n    return normalized\n', 'modeling_rwkv7': 'from __future__ import annotations\n\nfrom typing import Any, cast\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom transformers import GenerationMixin, PreTrainedModel\nfrom transformers.cache_utils import Cache\nfrom transformers.modeling_outputs import CausalLMOutputWithPast\n\nfrom .configuration_rwkv7 import RWKV7Config\nfrom .custom_ops import tilelang_state_finalize_op\nfrom .kernel_dispatch import resolve_backend  # type: ignore[reportMissingImports]\nfrom .kernel import state as _kernel_state\ntilelang_state_update = _kernel_state.tilelang_state_update  # type: ignore[reportMissingImports]\nfrom .state import RWKV7LayerState, RWKV7State  # type: ignore[reportMissingImports]\n\nTORCH_STACK = torch.stack  # type: ignore[attr-defined]\nTORCH_WHERE = torch.where  # type: ignore[attr-defined]\nFLOAT32 = torch.float32  # type: ignore[attr-defined]\nFLOAT16 = torch.float16  # type: ignore[attr-defined]\nBFLOAT16 = torch.bfloat16  # type: ignore[attr-defined]\nBOOL = torch.bool  # type: ignore[attr-defined]\nTORCH_EMPTY = torch.empty  # type: ignore[attr-defined]\nTORCH_CAT = torch.__dict__["cat"]\nTORCH_ZEROS_LIKE = torch.__dict__["zeros_like"]\nIS_GRAD_ENABLED = torch.is_grad_enabled  # type: ignore[attr-defined]\nTORCH_COMPILE = torch.compile  # type: ignore[attr-defined]\nIS_COMPILING = torch.compiler.is_compiling\n\n\n\ndef _cache_tensor_values(value: Any) -> list[torch.Tensor]:\n    if isinstance(value, torch.Tensor):\n        return [value]\n    if isinstance(value, (list, tuple)):\n        tensors: list[torch.Tensor] = []\n        for item in value:\n            tensors.extend(_cache_tensor_values(item))\n        return tensors\n    if value is not None and value.__class__.__name__ == "TileLangDecodeWorkspace":\n        tensors = []\n        for item in vars(value).values():\n            tensors.extend(_cache_tensor_values(item))\n        return tensors\n    return []\n\n\ndef _unique_storage_bytes(values: list[Any]) -> int:\n    storages: dict[tuple[str, int], int] = {}\n    for value in values:\n        for tensor in _cache_tensor_values(value):\n            storage = tensor.untyped_storage()\n            key = (str(tensor.device), storage.data_ptr())\n            storages[key] = max(storages.get(key, 0), storage.nbytes())\n    return sum(storages.values())\n\ndef _storage_keys(values: list[Any]) -> set[tuple[str, int]]:\n    keys = set()\n    for value in values:\n        for tensor in _cache_tensor_values(value):\n            storage = tensor.untyped_storage()\n            keys.add((str(tensor.device), storage.data_ptr()))\n    return keys\n\n\ndef _unique_storage_bytes_excluding(\n    values: list[Any], excluded: set[tuple[str, int]]\n) -> int:\n    storages: dict[tuple[str, int], int] = {}\n    for value in values:\n        for tensor in _cache_tensor_values(value):\n            storage = tensor.untyped_storage()\n            key = (str(tensor.device), storage.data_ptr())\n            if key not in excluded:\n                storages[key] = max(storages.get(key, 0), storage.nbytes())\n    return sum(storages.values())\n\n\ndef _replace_parameter_storage(\n    parameter: nn.Parameter, tensor: torch.Tensor\n) -> None:\n    with torch.no_grad():\n        parameter.set_(tensor)\n\n\ndef _restore_parameter_storage(parameter: nn.Parameter) -> None:\n    _replace_parameter_storage(\n        parameter, parameter.detach().clone(memory_format=torch.contiguous_format)\n    )\n\n\ndef _reset_elapsed_tokens(\n    elapsed_tokens: torch.Tensor | None,\n    reset: torch.Tensor | None,\n) -> None:\n    """Reset per-request token counts before processing a reset position."""\n    if elapsed_tokens is None or reset is None:\n        return\n    elapsed_tokens.masked_fill_(reset.reshape_as(elapsed_tokens), 0)\n\n\ndef _advance_elapsed_tokens(\n    elapsed_tokens: torch.Tensor | None,\n    active: torch.Tensor | None,\n) -> None:\n    """Advance each request once for an active token."""\n    if elapsed_tokens is None:\n        return\n    if active is None:\n        elapsed_tokens.add_(1)\n    else:\n        elapsed_tokens.add_(\n            active.reshape_as(elapsed_tokens).to(dtype=elapsed_tokens.dtype)\n        )\n\n\ndef _sequence_previous(\n    x: torch.Tensor,\n    initial: torch.Tensor,\n    active: torch.Tensor | None,\n    reset: torch.Tensor | None,\n) -> tuple[torch.Tensor, torch.Tensor]:\n    """Return previous active value for every sequence position."""\n    sequence_length = x.shape[1]\n    if active is None and reset is None:\n        previous = TORCH_CAT((initial.unsqueeze(1), x[:, :-1]), dim=1)\n        return previous, x[:, -1]\n\n    previous_value = initial\n    previous_steps: list[torch.Tensor] = []\n    for token_index in range(sequence_length):\n        if reset is not None:\n            reset_token = reset[:, token_index]\n            previous_value = TORCH_WHERE(\n                reset_token, TORCH_ZEROS_LIKE(previous_value), previous_value\n            )\n        previous_steps.append(previous_value)\n        current = x[:, token_index]\n        if active is None:\n            previous_value = current\n        else:\n            previous_value = TORCH_WHERE(\n                active[:, token_index], current, previous_value\n            )\n    return TORCH_STACK(previous_steps, dim=1), previous_value\n\n\ndef _sequence_matmul(\n    x: torch.Tensor, weight: torch.Tensor, tokenwise: bool\n) -> torch.Tensor:\n    if not tokenwise:\n        return x @ weight\n    return TORCH_STACK(\n        [x[:, token_index] @ weight for token_index in range(x.shape[1])],\n        dim=1,\n    )\n\n\ndef _sequence_linear(\n    module: nn.Linear, x: torch.Tensor, tokenwise: bool\n) -> torch.Tensor:\n    if not tokenwise:\n        return module(x)\n    return TORCH_STACK(\n        [module(x[:, token_index]) for token_index in range(x.shape[1])],\n        dim=1,\n    )\n\n\ndef _reset_layer_state(\n    state: RWKV7LayerState, reset: torch.Tensor | None\n) -> RWKV7LayerState:\n    if reset is None:\n        return state\n    matrix_reset = reset.unsqueeze(-1).unsqueeze(-1)\n    return RWKV7LayerState(\n        TORCH_WHERE(reset, TORCH_ZEROS_LIKE(state.channel), state.channel),\n        TORCH_WHERE(reset, TORCH_ZEROS_LIKE(state.time_shift), state.time_shift),\n        TORCH_WHERE(\n            matrix_reset, TORCH_ZEROS_LIKE(state.time_matrix), state.time_matrix\n        ),\n    )\n\nclass RWKV7Attention(nn.Module):\n    def __init__(self, config: RWKV7Config, layer_id: int) -> None:\n        super().__init__()\n        hidden = config.hidden_size\n        self.layer_id = layer_id\n        self.num_heads = config.num_attention_heads\n        self.head_size = config.head_size\n        self.kernel_backend = config.kernel_backend\n        self.hidden_size = hidden\n        self.intermediate_size = config.intermediate_size\n\n        self.x_r = nn.Parameter(TORCH_EMPTY(hidden))\n        self.x_w = nn.Parameter(TORCH_EMPTY(hidden))\n        self.x_k = nn.Parameter(TORCH_EMPTY(hidden))\n        self.x_v = nn.Parameter(TORCH_EMPTY(hidden))\n        self.x_a = nn.Parameter(TORCH_EMPTY(hidden))\n        self.x_g = nn.Parameter(TORCH_EMPTY(hidden))\n        self.w0 = nn.Parameter(TORCH_EMPTY(hidden))\n        self.w1 = nn.Parameter(TORCH_EMPTY(hidden, config.decay_lora_rank))\n        self.w2 = nn.Parameter(TORCH_EMPTY(config.decay_lora_rank, hidden))\n        self.a0 = nn.Parameter(TORCH_EMPTY(hidden))\n        self.a1 = nn.Parameter(TORCH_EMPTY(hidden, config.a_lora_rank))\n        self.a2 = nn.Parameter(TORCH_EMPTY(config.a_lora_rank, hidden))\n        self.g1 = nn.Parameter(TORCH_EMPTY(hidden, config.gate_lora_rank))\n        self.g2 = nn.Parameter(TORCH_EMPTY(config.gate_lora_rank, hidden))\n        self.k_k = nn.Parameter(TORCH_EMPTY(hidden))\n        self.k_a = nn.Parameter(TORCH_EMPTY(hidden))\n        self.r_k = nn.Parameter(TORCH_EMPTY(self.num_heads, self.head_size))\n\n        self.v0 = nn.Parameter(TORCH_EMPTY(hidden))\n        self.v1 = nn.Parameter(TORCH_EMPTY(hidden, config.value_lora_rank))\n        self.v2 = nn.Parameter(TORCH_EMPTY(config.value_lora_rank, hidden))\n        if layer_id == 0:\n            self.v0.requires_grad_(False)\n            self.v1.requires_grad_(False)\n            self.v2.requires_grad_(False)\n\n        self.receptance = nn.Linear(hidden, hidden, bias=False)\n        self.key = nn.Linear(hidden, hidden, bias=False)\n        self.value = nn.Linear(hidden, hidden, bias=False)\n        self.output = nn.Linear(hidden, hidden, bias=False)\n        self.ln_x = nn.GroupNorm(\n            self.num_heads,\n            hidden,\n            eps=config.head_size * config.layer_norm_epsilon,\n            affine=True,\n        )\n        self._x_mix_cache: torch.Tensor | None = None\n        self._x_mix_cache_versions: tuple[int, ...] | None = None\n        self.cuda_graph_time_mix = False\n        self._time_mix_graph: Any | None = None\n        self._time_mix_graph_key: tuple[Any, ...] | None = None\n        self._decode_backend: Any | None = None\n        self._decode_workspace: Any | None = None\n        self._decode_weight: torch.Tensor | None = None\n        self._decode_weight_key: tuple[Any, ...] | None = None\n        self._decode_rankout_weights: tuple[torch.Tensor, ...] | None = None\n        self._rkv_bmm_weight: torch.Tensor | None = None\n        self._rkv_bmm_weight_key: tuple[Any, ...] | None = None\n\n    def _x_mix_weights(self, x: torch.Tensor) -> torch.Tensor:\n        parameters = (self.x_r, self.x_w, self.x_k, self.x_v, self.x_a, self.x_g)\n        if self.training or IS_GRAD_ENABLED():\n            return TORCH_STACK(parameters)\n        if IS_COMPILING() and self._x_mix_cache is not None:\n            return self._x_mix_cache\n        versions = tuple(getattr(parameter, "_version", 0) for parameter in parameters)\n        cache = self._x_mix_cache\n        if (\n            cache is None\n            or cache.device != x.device\n            or cache.dtype != x.dtype\n            or self._x_mix_cache_versions != versions\n        ):\n            cache = TORCH_STACK(parameters).detach()\n            self._x_mix_cache = cache\n            self._x_mix_cache_versions = versions\n        return cache\n\n    def _packed_rkv_bmm_weight(\n        self, reference: torch.Tensor\n    ) -> torch.Tensor | None:\n        if (\n            self.kernel_backend != "tilelang"\n            or self.training\n            or IS_GRAD_ENABLED()\n            or reference.shape != (1, self.hidden_size)\n            or reference.dtype != FLOAT16\n            or reference.device.type != "cuda"\n            or torch.cuda.get_device_capability(reference.device) != (12, 0)\n        ):\n            return None\n        parameters = (\n            self.receptance.weight, self.key.weight, self.value.weight\n        )\n        cache_key = (\n            reference.device,\n            reference.dtype,\n            *(getattr(parameter, "_version", 0) for parameter in parameters),\n        )\n        if (\n            self._rkv_bmm_weight is None\n            or self._rkv_bmm_weight_key != cache_key\n        ):\n            self._rkv_bmm_weight = TORCH_STACK(\n                tuple(parameter.t() for parameter in parameters)\n            ).contiguous()\n            self._rkv_bmm_weight_key = cache_key\n        return self._rkv_bmm_weight\n\n    def _tilelang_backend_cache(self, reference: torch.Tensor) -> tuple[Any, Any]:\n        if (\n            self.kernel_backend != "tilelang"\n            or self.training\n            or IS_GRAD_ENABLED()\n            or reference.shape != (1, self.hidden_size)\n            or reference.dtype not in {torch.float16, torch.bfloat16}\n            or reference.device.type != "cuda"\n            or torch.cuda.get_device_capability(reference.device) != (12, 0)\n        ):\n            raise RuntimeError("TileLang B1T1 backend is unavailable")\n        from .tilelang_decode import TileLangDecodeBackend, TileLangDecodeSpec\n\n        backend = self._decode_backend\n        if (\n            backend is None\n            or self._decode_workspace is None\n            or backend.device != reference.device\n            or backend.dtype != reference.dtype\n        ):\n            spec = TileLangDecodeSpec(\n                channels=self.hidden_size,\n                ffn_rows=self.intermediate_size,\n                num_heads=self.num_heads,\n                head_size=self.head_size,\n                ranks=(\n                    self.w1.size(1),\n                    self.a1.size(1),\n                    self.g1.size(1),\n                    self.v1.size(1),\n                ),\n            )\n            backend = TileLangDecodeBackend(\n                spec, reference.device, reference.dtype\n            )\n            self._decode_backend = backend\n            self._decode_workspace = backend.create_workspace()\n            self._decode_weight = None\n            self._decode_rankout_weights = None\n            self._decode_weight_key = None\n        return backend, self._decode_workspace\n\n    def _tilelang_layernorm_mix6(\n        self,\n        residual: torch.Tensor,\n        previous: torch.Tensor,\n        norm: nn.LayerNorm,\n        active: torch.Tensor | None,\n    ) -> tuple[torch.Tensor, tuple[torch.Tensor, ...]] | None:\n        if (\n            self.kernel_backend != "tilelang"\n            or self.training\n            or IS_GRAD_ENABLED()\n            or residual.shape != (1, self.hidden_size)\n            or residual.dtype not in {torch.float16, torch.bfloat16}\n            or residual.device.type != "cuda"\n            or torch.cuda.get_device_capability(residual.device) != (12, 0)\n            or active is not None\n            or norm.weight is None\n            or norm.bias is None\n        ):\n            return None\n        backend, workspace = self._tilelang_backend_cache(residual)\n        normalized, mixed = backend.tmix_layernorm_mix6(\n            residual.view(-1),\n            previous.view(-1),\n            norm.weight,\n            norm.bias,\n            self._x_mix_weights(residual),\n            float(norm.eps),\n            workspace,\n        )\n        return (\n            normalized.view(1, self.hidden_size),\n            tuple(\n                value.view(1, self.hidden_size)\n                for value in mixed.unbind(dim=0)\n            ),\n        )\n\n    def _clear_inference_cache(self) -> None:\n        packed_weight = self._decode_weight\n        if packed_weight is not None:\n            packed_keys = _storage_keys([packed_weight])\n            for parameter in (\n                self.receptance.weight,\n                self.key.weight,\n                self.value.weight,\n                self.w1,\n                self.a1,\n                self.g1,\n                self.v1,\n            ):\n                if _storage_keys([parameter]) & packed_keys:\n                    _restore_parameter_storage(parameter)\n        rankout_weights = self._decode_rankout_weights\n        if rankout_weights is not None:\n            rankout_keys = _storage_keys([rankout_weights])\n            for parameter in (self.w2, self.a2, self.g2, self.v2):\n                if _storage_keys([parameter]) & rankout_keys:\n                    _restore_parameter_storage(parameter)\n        self._decode_backend = None\n        self._decode_workspace = None\n        self._decode_weight = None\n        self._decode_rankout_weights = None\n        self._decode_weight_key = None\n\n    def _prepare_tilelang_cache(\n        self, reference: torch.Tensor\n    ) -> tuple[Any, Any, torch.Tensor, tuple[torch.Tensor, ...], bool]:\n        from .tilelang_decode import TileLangDecodeBackend, TileLangDecodeSpec\n\n        parameters = (\n            self.receptance.weight,\n            self.key.weight,\n            self.value.weight,\n            self.w1,\n            self.a1,\n            self.g1,\n            self.v1,\n            self.w2,\n            self.a2,\n            self.g2,\n            self.v2,\n        )\n        weight_key = (\n            reference.device,\n            reference.dtype,\n            *(getattr(parameter, "_version", 0) for parameter in parameters),\n        )\n        created = (\n            self._decode_backend is None\n            or self._decode_workspace is None\n            or self._decode_weight is None\n            or self._decode_rankout_weights is None\n            or self._decode_weight_key != weight_key\n        )\n        if created:\n            self._clear_inference_cache()\n            spec = TileLangDecodeSpec(\n                channels=self.hidden_size,\n                ffn_rows=self.intermediate_size,\n                num_heads=self.num_heads,\n                head_size=self.head_size,\n                ranks=(\n                    self.w1.size(1),\n                    self.a1.size(1),\n                    self.g1.size(1),\n                    self.v1.size(1),\n                ),\n            )\n            backend = TileLangDecodeBackend(\n                spec, reference.device, reference.dtype\n            )\n            dense_parameters = (\n                self.receptance.weight, self.key.weight, self.value.weight\n            )\n            lowrank_parameters = (self.w1, self.a1, self.g1, self.v1)\n            packed_weight = backend.pack_rkv_weights(\n                dense_parameters,\n                tuple(\n                    parameter.t().contiguous()\n                    for parameter in lowrank_parameters\n                ),\n            )\n            rankout_parameters = (self.w2, self.a2, self.g2, self.v2)\n            rankout_weights = tuple(\n                parameter.t().contiguous()\n                for parameter in rankout_parameters\n            )\n            start = 0\n            for parameter in dense_parameters:\n                rows = parameter.size(0)\n                _replace_parameter_storage(\n                    parameter, packed_weight.narrow(0, start, rows)\n                )\n                start += rows\n            for parameter, rank in zip(\n                lowrank_parameters, spec.ranks, strict=True\n            ):\n                _replace_parameter_storage(\n                    parameter, packed_weight.narrow(0, start, rank).t()\n                )\n                start += rank\n            for parameter, packed in zip(\n                rankout_parameters, rankout_weights, strict=True\n            ):\n                _replace_parameter_storage(parameter, packed.t())\n            self._decode_backend = backend\n            self._decode_workspace = backend.create_workspace()\n            self._decode_weight = packed_weight\n            self._decode_rankout_weights = rankout_weights\n            self._decode_weight_key = (\n                reference.device,\n                reference.dtype,\n                *(\n                    getattr(parameter, "_version", 0)\n                    for parameter in parameters\n                ),\n            )\n        backend = self._decode_backend\n        workspace = self._decode_workspace\n        packed_weight = self._decode_weight\n        rankout_weights = self._decode_rankout_weights\n        if (\n            backend is None\n            or workspace is None\n            or packed_weight is None\n            or rankout_weights is None\n        ):\n            raise RuntimeError("TileLang decode cache initialization failed")\n        return backend, workspace, packed_weight, rankout_weights, created\n\n    def _tilelang_projections(\n        self,\n        xr: torch.Tensor,\n        xw: torch.Tensor,\n        xk: torch.Tensor,\n        xv: torch.Tensor,\n        xa: torch.Tensor,\n        xg: torch.Tensor,\n        first_value: torch.Tensor,\n    ) -> tuple[torch.Tensor, ...] | None:\n        if (\n            self.kernel_backend != "tilelang"\n            or self.training\n            or IS_GRAD_ENABLED()\n            or xr.shape != (1, self.hidden_size)\n            or xr.dtype not in {torch.float16, torch.bfloat16}\n            or xr.device.type != "cuda"\n            or torch.cuda.get_device_capability(xr.device) != (12, 0)\n        ):\n            return None\n        backend, workspace, packed_weight, rankout_weights, _ = (\n            self._prepare_tilelang_cache(xr)\n        )\n        backend.rkv(\n            tuple(\n                value.view(-1) for value in (xr, xk, xv, xw, xa, xg)\n            ),\n            packed_weight,\n            workspace.rkv_output,\n        )\n        (\n            receptance,\n            key,\n            value_base,\n            decay_rank,\n            a_rank,\n            g_rank,\n            value_rank,\n        ) = backend.rkv_views(workspace.rkv_output)\n        next_first_value = (\n            value_base.view(1, -1)\n            if self.layer_id == 0\n            else first_value\n        )\n        decay, gate_a, gate_g, value = backend.rankout_reduced(\n            (decay_rank, a_rank, g_rank, value_rank),\n            rankout_weights,\n            (self.a0, self.v0),\n            value_base,\n            next_first_value.view(-1),\n            workspace,\n            use_value_mix=self.layer_id != 0,\n        )\n        return (\n            decay.view(1, -1),\n            receptance.view(1, -1),\n            key.view(1, -1),\n            value.view(1, -1),\n            gate_a.view(1, -1),\n            gate_g.view(1, -1),\n            next_first_value,\n        )\n\n\n\n    def set_cuda_graph_time_mix(self, enabled: bool) -> None:\n        self.cuda_graph_time_mix = enabled\n        if not enabled:\n            self._time_mix_graph = None\n            self._time_mix_graph_key = None\n\n\n    def _cuda_graph_time_mix(\n        self,\n        x: torch.Tensor,\n        previous: torch.Tensor,\n        first_value: torch.Tensor,\n    ) -> tuple[torch.Tensor, ...] | None:\n        if (\n            not self.cuda_graph_time_mix\n            or self.training\n            or IS_GRAD_ENABLED()\n            or x.device.type != "cuda"\n            or x.shape[0] != 1\n        ):\n            return None\n        versions = tuple(\n            getattr(parameter, "_version", 0) for parameter in self.parameters()\n        )\n        key = (str(x.device), x.dtype, tuple(x.shape), versions)\n        if self._time_mix_graph is None or self._time_mix_graph_key != key:\n\n            def graph_function(\n                graph_x: torch.Tensor,\n                graph_previous: torch.Tensor,\n                graph_first_value: torch.Tensor,\n            ) -> tuple[torch.Tensor, ...]:\n                delta = graph_previous - graph_x\n                x_mix = TORCH_STACK(\n                    (self.x_r, self.x_w, self.x_k, self.x_v, self.x_a, self.x_g)\n                )\n                xr, xw, xk, xv, xa, xg = (\n                    graph_x.unsqueeze(1) + delta.unsqueeze(1) * x_mix\n                ).unbind(dim=1)\n                decay = self.w0 + (xw @ self.w1).tanh() @ self.w2\n                decay = (-0.606531 * decay.sigmoid()).exp()\n                receptance = self.receptance(xr)\n                key_tensor = self.key(xk)\n                value = self.value(xv)\n                if self.layer_id == 0:\n                    graph_first_value = value\n                else:\n                    value = value + (graph_first_value - value) * (\n                        self.v0 + (xv @ self.v1) @ self.v2\n                    ).sigmoid()\n                gate_a = (self.a0 + (xa @ self.a1) @ self.a2).sigmoid()\n                gate_g = (xg @ self.g1).sigmoid() @ self.g2\n                return (\n                    decay,\n                    receptance,\n                    key_tensor,\n                    value,\n                    gate_a,\n                    gate_g,\n                    graph_first_value,\n                )\n\n            self._time_mix_graph = TORCH_COMPILE(\n                graph_function,\n                backend="cudagraphs",\n                fullgraph=True,\n                dynamic=False,\n            )\n            self._time_mix_graph_key = key\n        return self._time_mix_graph(x, previous, first_value)\n\n    def forward(\n        self,\n        x: torch.Tensor,\n        state: RWKV7LayerState,\n        first_value: torch.Tensor,\n        active: torch.Tensor | None,\n        elapsed_tokens: torch.Tensor | None = None,\n        mixed_inputs: tuple[torch.Tensor, ...] | None = None,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n        batch_size = x.shape[0]\n        heads = self.num_heads\n        head_size = self.head_size\n        decay_delta: torch.Tensor | None = None\n\n        graph_outputs = (\n            None\n            if mixed_inputs is not None\n            else self._cuda_graph_time_mix(x, state.time_shift, first_value)\n        )\n        if graph_outputs is None:\n            if mixed_inputs is None:\n                delta = state.time_shift - x\n                x_mix = self._x_mix_weights(x)\n                xr, xw, xk, xv, xa, xg = (\n                    x.unsqueeze(1) + delta.unsqueeze(1) * x_mix\n                ).unbind(dim=1)\n            else:\n                xr, xw, xk, xv, xa, xg = mixed_inputs\n\n            projections = self._tilelang_projections(\n                xr, xw, xk, xv, xa, xg, first_value\n            )\n            if projections is None:\n                decay_delta = (xw @ self.w1).tanh() @ self.w2\n                decay = self.w0 + decay_delta\n                rkv_weight = self._packed_rkv_bmm_weight(xr)\n                if rkv_weight is None:\n                    receptance = self.receptance(xr)\n                    key = self.key(xk)\n                    value = self.value(xv)\n                else:\n                    rkv = torch.bmm(\n                        TORCH_STACK((xr, xk, xv), dim=0), rkv_weight\n                    )\n                    receptance, key, value = rkv.unbind(dim=0)\n                value_rank = xv @ self.v1\n                a_rank = xa @ self.a1\n                g_rank = xg @ self.g1\n                if self.layer_id == 0:\n                    first_value = value\n                else:\n                    value = value + (first_value - value) * (\n                        self.v0 + value_rank @ self.v2\n                    ).sigmoid()\n                gate_a = (self.a0 + a_rank @ self.a2).sigmoid()\n                gate_g = g_rank.sigmoid() @ self.g2\n            else:\n                (\n                    decay_delta,\n                    receptance,\n                    key,\n                    value,\n                    gate_a,\n                    gate_g,\n                    first_value,\n                ) = projections\n                decay = self.w0 + decay_delta\n        else:\n            decay, receptance, key, value, gate_a, gate_g, first_value = (\n                graph_outputs\n            )\n\n        receptance = receptance.view(batch_size, heads, head_size, 1)\n        value = value.view(batch_size, heads, head_size, 1)\n\n        key_gate_outputs = None\n        if (\n            self.kernel_backend == "tilelang"\n            and not self.training\n            and not IS_GRAD_ENABLED()\n            and batch_size == 1\n            and x.dtype in {FLOAT16, BFLOAT16}\n            and x.device.type == "cuda"\n            and torch.cuda.get_device_capability(x.device) == (12, 0)\n        ):\n            key_backend, key_workspace = self._tilelang_backend_cache(x)\n            key_gate_outputs = key_backend.key_gate(\n                key.view(-1),\n                self.k_k,\n                gate_a.view(-1),\n                self.k_a,\n                key_workspace,\n            )\n        if key_gate_outputs is None:\n            normalized_key = key * self.k_k\n            normalized_key = F.normalize(\n                normalized_key.view(batch_size, heads, head_size), dim=-1\n            ).view(batch_size, -1)\n            key = (key * (1 + (gate_a - 1) * self.k_a)).view(\n                batch_size, heads, 1, head_size\n            )\n        else:\n            normalized_key, modified_key, _, _ = key_gate_outputs\n            normalized_key = normalized_key.view(batch_size, -1)\n            key = modified_key.view(batch_size, heads, 1, head_size)\n\n        tile_mixed = None\n        next_matrix = state.time_matrix\n        backend = "torch"\n        decode_backend = None\n        decode_workspace = None\n        training_tilelang = (\n            self.training\n            and self.kernel_backend == "tilelang"\n            and x.device.type == "cuda"\n            and state.time_matrix.dtype == FLOAT32\n        )\n        if training_tilelang:\n            backend = "tilelang"\n        elif not self.training and self.kernel_backend != "torch":\n            backend = resolve_backend(self.kernel_backend, x.device)\n        if backend == "tilelang" and state.time_matrix.dtype not in {\n            FLOAT16,\n            BFLOAT16,\n            FLOAT32,\n        }:\n            backend = "torch"\n        # FP32 state uses exact TileLang pointwise finalization followed by the\n        # native PyTorch projection. SM120 never selects the fused projection.\n        if backend == "tilelang":\n            try:\n                if batch_size == 1 and not training_tilelang:\n                    decode_backend, decode_workspace = (\n                        self._tilelang_backend_cache(x)\n                    )\n                use_fused_wkv = (\n                    batch_size == 1\n                    and x.dtype == FLOAT16\n                    and state.time_matrix.dtype == FLOAT16\n                    and decay_delta is not None\n                    and elapsed_tokens is not None\n                    and active is None\n                    and decode_backend is not None\n                    and decode_workspace is not None\n                    and key_gate_outputs is not None\n                )\n                if use_fused_wkv:\n                    decode_backend.wkv_w0_t1(\n                        state.time_matrix,\n                        receptance.squeeze(-1).view(\n                            1, 1, heads, head_size\n                        ),\n                        decay_delta.view(1, 1, heads, head_size),\n                        self.w0.view(heads, head_size),\n                        key.squeeze(-2).view(1, 1, heads, head_size),\n                        value.squeeze(-1).view(1, 1, heads, head_size),\n                        decode_workspace.anti_key.view(\n                            1, 1, heads, head_size\n                        ),\n                        decode_workspace.anti_gate.view(\n                            1, 1, heads, head_size\n                        ),\n                        elapsed_tokens,\n                        decode_workspace.wkv_output,\n                    )\n                    next_matrix = state.time_matrix\n                    tile_mixed = decode_workspace.wkv_output.view(\n                        1, heads, head_size\n                    )\n                else:\n                    if decay_delta is not None:\n                        decay = (-0.606531 * decay.sigmoid()).exp()\n                    next_matrix, tile_mixed = tilelang_state_update(\n                        state.time_matrix,\n                        decay.view(batch_size, heads, head_size),\n                        normalized_key.view(batch_size, heads, head_size),\n                        gate_a.view(batch_size, heads, head_size),\n                        value.squeeze(-1),\n                        key.squeeze(-2),\n                        receptance.squeeze(-1),\n                        state_finalize_op=tilelang_state_finalize_op\n                        if training_tilelang\n                        else None,\n                    )\n            except Exception:\n                if self.kernel_backend != "auto":\n                    raise\n                backend = "torch"\n        if backend == "torch":\n            if decay_delta is not None:\n                decay = (-0.606531 * decay.sigmoid()).exp()\n            decay = decay.view(batch_size, heads, 1, head_size)\n            value_key = value @ key\n            anti_value_key = (-normalized_key).view(\n                batch_size, heads, head_size, 1\n            ) @ (normalized_key * gate_a).view(\n                batch_size, heads, 1, head_size\n            )\n            matrix_f32 = state.time_matrix.float()\n            next_matrix = matrix_f32 * decay.to(FLOAT32)\n            next_matrix = (\n                next_matrix + matrix_f32 @ anti_value_key.to(FLOAT32)\n            )\n            next_matrix = next_matrix + value_key.to(FLOAT32)\n            if state.time_matrix.dtype != FLOAT32:\n                next_matrix = next_matrix.to(state.time_matrix.dtype)\n        next_shift = x\n        if active is not None:\n            next_shift = TORCH_WHERE(active, next_shift, state.time_shift)\n            next_matrix = TORCH_WHERE(\n                active.unsqueeze(-1).unsqueeze(-1),\n                next_matrix,\n                state.time_matrix,\n            )\n\n        projected = (\n            tile_mixed\n            if tile_mixed is not None\n            else (next_matrix.to(dtype=x.dtype) @ receptance).squeeze(-1)\n        )\n        decode_backend = self._decode_backend\n        decode_workspace = self._decode_workspace\n        if (\n            self.kernel_backend == "tilelang"\n            and batch_size == 1\n            and decode_backend is not None\n            and decode_workspace is not None\n            and self.ln_x.weight is not None\n            and self.ln_x.bias is not None\n        ):\n            mixed = decode_backend.post_state(\n                projected.view(heads, head_size),\n                receptance.squeeze(-1).view(heads, head_size),\n                key.squeeze(-2).view(heads, head_size),\n                value.squeeze(-1).view(heads, head_size),\n                self.r_k,\n                gate_g.view(-1),\n                self.ln_x.weight,\n                self.ln_x.bias,\n                float(self.ln_x.eps),\n                decode_workspace,\n            ).view(1, heads * head_size)\n        else:\n            mixed = self.ln_x(projected.flatten(start_dim=1))\n            rkv = (\n                receptance.squeeze(-1) * key.squeeze(-2) * self.r_k\n            ).sum(dim=-1, keepdim=True) * value.squeeze(-1)\n            mixed = (\n                mixed + rkv.view(batch_size, heads * head_size)\n            ) * gate_g\n        if (\n            self.kernel_backend == "tilelang"\n            and batch_size == 1\n            and decode_backend is not None\n            and decode_workspace is not None\n        ):\n            decode_backend.gemv(\n                mixed.view(-1),\n                self.output.weight,\n                decode_workspace.attention_output,\n            )\n            attention_output = decode_workspace.attention_output.view(\n                1, self.hidden_size\n            )\n        else:\n            attention_output = self.output(mixed)\n        return attention_output, first_value, next_shift, next_matrix\n\n    def forward_sequence(\n        self,\n        x: torch.Tensor,\n        state: RWKV7LayerState,\n        first_value: torch.Tensor,\n        active: torch.Tensor | None,\n        reset: torch.Tensor | None,\n        tokenwise_projections: bool = False,\n        state_scan_backend: str = "torch",\n        elapsed_tokens: torch.Tensor | None = None,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n        batch_size, sequence_length, _ = x.shape\n        heads = self.num_heads\n        head_size = self.head_size\n\n        previous, next_shift = _sequence_previous(\n            x, state.time_shift, active, reset\n        )\n        delta = previous - x\n        x_mix = self._x_mix_weights(x)\n        xr, xw, xk, xv, xa, xg = (\n            x.unsqueeze(2) + delta.unsqueeze(2) * x_mix\n        ).unbind(dim=2)\n\n        decay_raw = self.w0 + _sequence_matmul(\n            _sequence_matmul(xw, self.w1, tokenwise_projections).tanh(),\n            self.w2,\n            tokenwise_projections,\n        )\n        decay = (-0.606531 * decay_raw.sigmoid()).exp().view(\n            batch_size, sequence_length, heads, head_size\n        )\n        receptance = _sequence_linear(\n            self.receptance, xr, tokenwise_projections\n        ).view(batch_size, sequence_length, heads, head_size)\n        key = _sequence_linear(self.key, xk, tokenwise_projections)\n        value = _sequence_linear(self.value, xv, tokenwise_projections)\n        if self.layer_id == 0:\n            first_value = value\n        else:\n            value = value + (first_value - value) * (\n                self.v0\n                + _sequence_matmul(\n                    _sequence_matmul(xv, self.v1, tokenwise_projections),\n                    self.v2,\n                    tokenwise_projections,\n                )\n            ).sigmoid()\n\n        value = value.view(batch_size, sequence_length, heads, head_size)\n        gate_a = (\n            self.a0\n            + _sequence_matmul(\n                _sequence_matmul(xa, self.a1, tokenwise_projections),\n                self.a2,\n                tokenwise_projections,\n            )\n        ).sigmoid()\n        gate_g = _sequence_matmul(\n            _sequence_matmul(xg, self.g1, tokenwise_projections).sigmoid(),\n            self.g2,\n            tokenwise_projections,\n        )\n\n        normalized_key = key * self.k_k\n        normalized_key = F.normalize(\n            normalized_key.view(batch_size, sequence_length, heads, head_size),\n            dim=-1,\n        )\n        key = (key * (1 + (gate_a - 1) * self.k_a)).view(\n            batch_size, sequence_length, heads, head_size\n        )\n        gate_a = gate_a.view(batch_size, sequence_length, heads, head_size)\n\n        if state_scan_backend == "tilelang-wkv":\n            if active is not None or reset is not None:\n                raise RuntimeError(\n                    "TileLang WKV prefill does not support masks or resets"\n                )\n            if elapsed_tokens is None:\n                raise RuntimeError("TileLang WKV prefill requires elapsed-token state")\n            from .kernel import decode as _kernel_decode\n            _wkv_precise_out = _kernel_decode._wkv_precise_out\n\n            mixed = torch.empty_like(receptance)\n            _wkv_precise_out(\n                state.time_matrix,\n                receptance,\n                decay,\n                key,\n                value,\n                -normalized_key,\n                normalized_key * gate_a,\n                elapsed_tokens,\n                mixed,\n            )\n            matrix = state.time_matrix\n        elif state_scan_backend == "tilelang-fast":\n            if active is not None or reset is not None:\n                raise RuntimeError(\n                    "Fast TileLang sequence scan does not support masks or resets"\n                )\n            from .kernel import state as _kernel_state\n            tilelang_fast_state_scan = _kernel_state.tilelang_fast_state_scan  # type: ignore[reportMissingImports]\n\n            matrix, mixed = tilelang_fast_state_scan(\n                state.time_matrix,\n                decay,\n                normalized_key,\n                gate_a,\n                value,\n                key,\n                receptance,\n            )\n        elif state_scan_backend == "tilelang":\n            from .kernel import state as _kernel_state\n            tilelang_state_scan = _kernel_state.tilelang_state_scan  # type: ignore[reportMissingImports]\n\n            matrix, mixed = tilelang_state_scan(\n                state.time_matrix,\n                decay,\n                normalized_key,\n                gate_a,\n                value,\n                key,\n                receptance,\n                active=active,\n                reset=reset,\n                output_mode="full",\n            )\n        else:\n            matrix = state.time_matrix\n            mixed_steps: list[torch.Tensor] = []\n            for token_index in range(sequence_length):\n                if reset is not None:\n                    reset_token = reset[:, token_index].unsqueeze(-1).unsqueeze(-1)\n                    matrix = TORCH_WHERE(\n                        reset_token, TORCH_ZEROS_LIKE(matrix), matrix\n                    )\n                normalized_token = normalized_key[:, token_index]\n                anti_value_key = (-normalized_token).unsqueeze(-1) @ (\n                    normalized_token * gate_a[:, token_index]\n                ).unsqueeze(-2)\n                value_key = value[:, token_index].unsqueeze(-1) @ key[\n                    :, token_index\n                ].unsqueeze(-2)\n                matrix_f32 = matrix.float()\n                candidate = (\n                    matrix_f32\n                    * decay[:, token_index].to(FLOAT32).unsqueeze(-2)\n                )\n                candidate = (\n                    candidate + matrix_f32 @ anti_value_key.to(FLOAT32)\n                )\n                candidate = candidate + value_key.to(FLOAT32)\n                if matrix.dtype != FLOAT32:\n                    candidate = candidate.to(matrix.dtype)\n                if active is None:\n                    matrix = candidate\n                else:\n                    active_token = active[:, token_index].unsqueeze(-1).unsqueeze(-1)\n                    matrix = TORCH_WHERE(active_token, candidate, matrix)\n                mixed_steps.append(\n                    (\n                        matrix.to(dtype=x.dtype)\n                        @ receptance[:, token_index].unsqueeze(-1)\n                    ).squeeze(-1)\n                )\n            mixed = TORCH_STACK(mixed_steps, dim=1)\n        mixed = self.ln_x(mixed.reshape(batch_size * sequence_length, -1)).view(\n            batch_size, sequence_length, heads * head_size\n        )\n        rkv = (receptance * key * self.r_k).sum(dim=-1, keepdim=True) * value\n        mixed = (mixed + rkv.reshape(batch_size, sequence_length, -1)) * gate_g\n        return (\n            _sequence_linear(self.output, mixed, tokenwise_projections),\n            first_value,\n            next_shift,\n            matrix,\n        )\n\n\nclass RWKV7FeedForward(nn.Module):\n    def __init__(self, config: RWKV7Config) -> None:\n        super().__init__()\n        self.x_k = nn.Parameter(TORCH_EMPTY(config.hidden_size))\n        self.key = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)\n        self.value = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)\n        self.kernel_backend = config.kernel_backend\n        self.hidden_size = config.hidden_size\n        self.intermediate_size = config.intermediate_size\n        self.num_heads = config.num_attention_heads\n        self.head_size = config.head_size\n        self.ranks = (\n            config.decay_lora_rank,\n            config.a_lora_rank,\n            config.gate_lora_rank,\n            config.value_lora_rank,\n        )\n        self._decode_backend: Any | None = None\n        self._decode_workspace: Any | None = None\n        self._decode_value_weight: torch.Tensor | None = None\n        self._decode_weight_key: tuple[Any, ...] | None = None\n\n    def _clear_inference_cache(self) -> None:\n        packed_weight = self._decode_value_weight\n        if packed_weight is not None:\n            packed_keys = _storage_keys([packed_weight])\n            if _storage_keys([self.value.weight]) & packed_keys:\n                _restore_parameter_storage(self.value.weight)\n        self._decode_backend = None\n        self._decode_workspace = None\n        self._decode_value_weight = None\n        self._decode_weight_key = None\n\n    def _tilelang_supported(self, reference: torch.Tensor) -> bool:\n        return (\n            self.kernel_backend == "tilelang"\n            and not self.training\n            and not IS_GRAD_ENABLED()\n            and reference.shape == (1, self.hidden_size)\n            and reference.dtype in {torch.float16, torch.bfloat16}\n            and reference.device.type == "cuda"\n            and torch.cuda.get_device_capability(reference.device) == (12, 0)\n        )\n\n    def _tilelang_cache(\n        self, reference: torch.Tensor\n    ) -> tuple[Any, Any, torch.Tensor, bool]:\n        from .tilelang_decode import TileLangDecodeBackend, TileLangDecodeSpec\n\n        weight_key = (\n            reference.device,\n            reference.dtype,\n            getattr(self.key.weight, "_version", 0),\n            getattr(self.value.weight, "_version", 0),\n        )\n        created = (\n            self._decode_backend is None\n            or self._decode_workspace is None\n            or self._decode_value_weight is None\n            or self._decode_weight_key != weight_key\n        )\n        if created:\n            self._clear_inference_cache()\n            spec = TileLangDecodeSpec(\n                channels=self.hidden_size,\n                ffn_rows=self.intermediate_size,\n                num_heads=self.num_heads,\n                head_size=self.head_size,\n                ranks=self.ranks,\n            )\n            backend = TileLangDecodeBackend(\n                spec, reference.device, reference.dtype\n            )\n            packed_value_weight = backend.pack_ffn_value_weight(\n                self.value.weight\n            )\n            _replace_parameter_storage(\n                self.value.weight, packed_value_weight.t()\n            )\n            self._decode_backend = backend\n            self._decode_workspace = backend.create_workspace()\n            self._decode_value_weight = packed_value_weight\n            self._decode_weight_key = (\n                reference.device,\n                reference.dtype,\n                getattr(self.key.weight, "_version", 0),\n                getattr(self.value.weight, "_version", 0),\n            )\n        backend = self._decode_backend\n        workspace = self._decode_workspace\n        packed_value_weight = self._decode_value_weight\n        if (\n            backend is None\n            or workspace is None\n            or packed_value_weight is None\n        ):\n            raise RuntimeError("TileLang FFN cache initialization failed")\n        return backend, workspace, packed_value_weight, created\n\n    def _tilelang_output(self, mixed: torch.Tensor) -> torch.Tensor | None:\n        if not self._tilelang_supported(mixed):\n            return None\n        backend, workspace, packed_value_weight, _ = self._tilelang_cache(mixed)\n        backend.ffn(\n            mixed.view(-1),\n            self.key.weight,\n            packed_value_weight,\n            workspace.ffn_output,\n            workspace,\n        )\n        return workspace.ffn_output.view(1, self.hidden_size)\n\n    def _tilelang_output_from_residual(\n        self,\n        residual: torch.Tensor,\n        previous: torch.Tensor,\n        norm: nn.LayerNorm,\n    ) -> tuple[torch.Tensor, torch.Tensor] | None:\n        if (\n            not self._tilelang_supported(residual)\n            or norm.weight is None\n            or norm.bias is None\n        ):\n            return None\n        backend, workspace, packed_value_weight, _ = self._tilelang_cache(residual)\n        normalized, mixed = backend.cmix_layernorm_mix(\n            residual.view(-1),\n            previous.view(-1),\n            norm.weight,\n            norm.bias,\n            self.x_k,\n            float(norm.eps),\n            workspace,\n        )\n        backend.ffn(\n            mixed,\n            self.key.weight,\n            packed_value_weight,\n            workspace.ffn_output,\n            workspace,\n        )\n        return (\n            workspace.ffn_output.view(1, self.hidden_size),\n            normalized.view(1, self.hidden_size),\n        )\n\n    def _tilelang_output_from_residual_update(\n        self,\n        residual: torch.Tensor,\n        update: torch.Tensor,\n        previous: torch.Tensor,\n        norm: nn.LayerNorm,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None:\n        if (\n            not self._tilelang_supported(residual)\n            or update.shape != residual.shape\n            or norm.weight is None\n            or norm.bias is None\n        ):\n            return None\n        backend, workspace, packed_value_weight, _ = self._tilelang_cache(residual)\n        combined, normalized, mixed = backend.cmix_add_layernorm_mix(\n            residual.view(-1),\n            update.view(-1),\n            previous.view(-1),\n            norm.weight,\n            norm.bias,\n            self.x_k,\n            float(norm.eps),\n            workspace,\n        )\n        backend.ffn(\n            mixed,\n            self.key.weight,\n            packed_value_weight,\n            workspace.ffn_output,\n            workspace,\n            residual=combined.view(-1),\n        )\n        return (\n            workspace.ffn_output.view(1, self.hidden_size),\n            normalized.view(1, self.hidden_size),\n            combined.view(1, self.hidden_size),\n        )\n\n    def forward(\n        self,\n        x: torch.Tensor,\n        previous: torch.Tensor,\n        active: torch.Tensor | None,\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        delta = previous - x\n        key = x + delta * self.x_k\n        output = self._tilelang_output(key)\n        if output is None:\n            hidden = F.relu(self.key(key)).square()\n            output = self.value(hidden)\n        next_channel = x\n        if active is not None:\n            next_channel = TORCH_WHERE(active, next_channel, previous)\n        return output, next_channel\n\n    def forward_from_residual(\n        self,\n        residual: torch.Tensor,\n        previous: torch.Tensor,\n        norm: nn.LayerNorm,\n        active: torch.Tensor | None,\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        tilelang_output = self._tilelang_output_from_residual(\n            residual, previous, norm\n        )\n        if tilelang_output is None:\n            return self.forward(norm(residual), previous, active)\n        output, next_channel = tilelang_output\n        if active is not None:\n            next_channel = TORCH_WHERE(active, next_channel, previous)\n        return output, next_channel\n\n    def forward_from_residual_update(\n        self,\n        residual: torch.Tensor,\n        update: torch.Tensor,\n        previous: torch.Tensor,\n        norm: nn.LayerNorm,\n        active: torch.Tensor | None,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n        tilelang_output = (\n            None\n            if active is not None\n            else self._tilelang_output_from_residual_update(\n                residual, update, previous, norm\n            )\n        )\n        if tilelang_output is None:\n            combined = residual + update\n            output, next_channel = self.forward_from_residual(\n                combined, previous, norm, active\n            )\n            return combined + output, next_channel, combined\n        output, next_channel, combined = tilelang_output\n        return output, next_channel, combined\n\n    def forward_sequence(\n        self,\n        x: torch.Tensor,\n        previous: torch.Tensor,\n        active: torch.Tensor | None,\n        reset: torch.Tensor | None,\n        tokenwise_projections: bool = False,\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        previous_sequence, next_channel = _sequence_previous(\n            x, previous, active, reset\n        )\n        key = x + (previous_sequence - x) * self.x_k\n        hidden = F.relu(\n            _sequence_linear(self.key, key, tokenwise_projections)\n        ).square()\n        output = _sequence_linear(self.value, hidden, tokenwise_projections)\n        return output, next_channel\n\n\nclass RWKV7Block(nn.Module):\n    def __init__(self, config: RWKV7Config, layer_id: int) -> None:\n        super().__init__()\n        hidden = config.hidden_size\n        if layer_id == 0:\n            self.ln0 = nn.LayerNorm(hidden, eps=config.layer_norm_epsilon)\n        self.ln1 = nn.LayerNorm(hidden, eps=config.layer_norm_epsilon)\n        self.ln2 = nn.LayerNorm(hidden, eps=config.layer_norm_epsilon)\n        self.att = RWKV7Attention(config, layer_id)\n        self.ffn = RWKV7FeedForward(config)\n\n    def forward(\n        self,\n        x: torch.Tensor,\n        state: RWKV7LayerState,\n        first_value: torch.Tensor,\n        active: torch.Tensor | None,\n        elapsed_tokens: torch.Tensor | None = None,\n    ) -> tuple[torch.Tensor, RWKV7LayerState, torch.Tensor]:\n        time_mix = self.att._tilelang_layernorm_mix6(\n            x, state.time_shift, self.ln1, active\n        )\n        if time_mix is None:\n            normalized = self.ln1(x)\n            mixed_inputs = None\n        else:\n            normalized, mixed_inputs = time_mix\n        mixed, first_value, time_shift, time_matrix = self.att(\n            normalized,\n            state,\n            first_value,\n            active,\n            elapsed_tokens,\n            mixed_inputs,\n        )\n        x, channel, _ = self.ffn.forward_from_residual_update(\n            x, mixed, state.channel, self.ln2, active\n        )\n        return x, RWKV7LayerState(channel, time_shift, time_matrix), first_value\n\n\n    def forward_sequence(\n        self,\n        x: torch.Tensor,\n        state: RWKV7LayerState,\n        first_value: torch.Tensor,\n        active: torch.Tensor | None,\n        reset: torch.Tensor | None,\n        tokenwise_projections: bool = False,\n        state_scan_backend: str = "torch",\n        elapsed_tokens: torch.Tensor | None = None,\n    ) -> tuple[torch.Tensor, RWKV7LayerState, torch.Tensor]:\n        mixed, first_value, time_shift, time_matrix = self.att.forward_sequence(\n            self.ln1(x),\n            state,\n            first_value,\n            active,\n            reset,\n            tokenwise_projections,\n            state_scan_backend,\n            elapsed_tokens,\n        )\n        x = x + mixed\n        mixed, channel = self.ffn.forward_sequence(\n            self.ln2(x),\n            state.channel,\n            active,\n            reset,\n            tokenwise_projections,\n        )\n        x = x + mixed\n        return x, RWKV7LayerState(channel, time_shift, time_matrix), first_value\n\n\nclass RWKV7ForCausalLM(PreTrainedModel, GenerationMixin):\n    config_class = RWKV7Config\n    base_model_prefix = "rwkv7"\n    main_input_name = "input_ids"\n    _is_stateful = True\n    _supports_cache_class = True\n    supports_gradient_checkpointing = True\n    _supports_sdpa = False\n    _supports_flash_attn = False\n\n    def __init__(self, config: RWKV7Config) -> None:\n        super().__init__(config)\n        self.emb = nn.Embedding(config.vocab_size, config.hidden_size)\n        self.blocks = nn.ModuleList(\n            [\n                RWKV7Block(config, layer_id)\n                for layer_id in range(config.num_hidden_layers)\n            ]\n        )\n        self.ln_out = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_epsilon)\n        self.head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)\n        self.kernel_backend = config.kernel_backend\n        self._inference_cache_epoch = 0\n        self.gradient_checkpointing = False\n        self.post_init()\n\n    def _init_weights(self, module: nn.Module) -> None:\n        if isinstance(module, (nn.Linear, nn.Embedding)):\n            nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)\n        elif isinstance(module, (nn.LayerNorm, nn.GroupNorm)):\n            if module.weight is not None:\n                nn.init.ones_(module.weight)\n            if module.bias is not None:\n                nn.init.zeros_(module.bias)\n        for name, parameter in module.named_parameters(recurse=False):\n            if name not in {"weight", "bias"}:\n                nn.init.zeros_(parameter)\n\n    def get_input_embeddings(self) -> nn.Embedding:\n        return self.emb\n\n    def set_input_embeddings(self, value: nn.Module) -> None:\n        if not isinstance(value, nn.Embedding):\n            raise TypeError("Input embeddings must be nn.Embedding")\n        self.emb = value\n\n    def get_output_embeddings(self) -> nn.Linear:\n        return self.head\n\n    def set_output_embeddings(self, new_embeddings: nn.Module) -> None:\n        if not isinstance(new_embeddings, nn.Linear):\n            raise TypeError("Output embeddings must be nn.Linear")\n        self.head = new_embeddings\n\n    @property\n    def inference_cache_epoch(self) -> int:\n        """Generation counter used to invalidate captured inference runners."""\n        return self._inference_cache_epoch\n\n    def inference_cache_stats(self) -> dict[str, int]:\n        """Return bounded tensor/cache counts without allocating new device storage."""\n        attention_values: list[Any] = []\n        mix_values: list[Any] = []\n        ffn_values: list[Any] = []\n        workspace_values: list[Any] = []\n        backend_count = 0\n        workspace_count = 0\n        time_mix_graph_count = 0\n        for block in cast(Any, self.blocks):\n            mix_values.append(block.att._x_mix_cache)\n            attention_values.extend(\n                (\n                    block.att._rkv_bmm_weight,\n                    block.att._decode_weight,\n                    block.att._decode_rankout_weights,\n                )\n            )\n            ffn_values.append(block.ffn._decode_value_weight)\n            workspace_values.extend(\n                (block.att._decode_workspace, block.ffn._decode_workspace)\n            )\n            backend_count += int(block.att._decode_backend is not None)\n            backend_count += int(block.ffn._decode_backend is not None)\n            workspace_count += int(block.att._decode_workspace is not None)\n            workspace_count += int(block.ffn._decode_workspace is not None)\n            time_mix_graph_count += int(block.att._time_mix_graph is not None)\n        packed_values = attention_values + ffn_values\n        parameter_values: list[Any] = list(self.parameters())\n        parameter_keys = _storage_keys(parameter_values)\n        packed_bytes = _unique_storage_bytes(packed_values)\n        extra_packed_bytes = _unique_storage_bytes_excluding(\n            packed_values, parameter_keys\n        )\n        workspace_bytes = _unique_storage_bytes(workspace_values)\n        all_cache_values = packed_values + mix_values + workspace_values\n        return {\n            "parameter_storage_bytes": _unique_storage_bytes(parameter_values),\n            "mix_cache_bytes": _unique_storage_bytes(mix_values),\n            "packed_attention_bytes": _unique_storage_bytes(attention_values),\n            "packed_ffn_bytes": _unique_storage_bytes(ffn_values),\n            "packed_weight_bytes": packed_bytes,\n            "shared_layout_bytes": packed_bytes - extra_packed_bytes,\n            "extra_packed_weight_bytes": extra_packed_bytes,\n            "workspace_bytes": workspace_bytes,\n            "extra_tensor_bytes": _unique_storage_bytes_excluding(\n                all_cache_values, parameter_keys\n            ),\n            "total_tensor_bytes": _unique_storage_bytes(\n                all_cache_values\n            ),\n            "backend_count": backend_count,\n            "workspace_count": workspace_count,\n            "time_mix_graph_count": time_mix_graph_count,\n            "epoch": self._inference_cache_epoch,\n        }\n\n    def clear_inference_caches(self, *, include_compiled: bool = False) -> None:\n        """Release model-owned inference layouts, workspaces, and graph wrappers.\n\n        Captured runners record the cache epoch and refuse replay after this call.\n        Global compiled TileLang kernels remain cached unless ``include_compiled`` is\n        requested explicitly.\n        """\n        for block in cast(Any, self.blocks):\n            attention = block.att\n            attention._x_mix_cache = None\n            attention._x_mix_cache_versions = None\n            attention._time_mix_graph = None\n            attention._time_mix_graph_key = None\n            attention._clear_inference_cache()\n            attention._rkv_bmm_weight = None\n            attention._rkv_bmm_weight_key = None\n            feed_forward = block.ffn\n            feed_forward._clear_inference_cache()\n        self._inference_cache_epoch += 1\n        if include_compiled:\n            from .kernel import decode as _kernel_decode\n            clear_tilelang_kernel_caches = _kernel_decode.clear_tilelang_kernel_caches\n            from .kernel import state as _kernel_state\n            clear_tilelang_state_kernel_caches = _kernel_state.clear_tilelang_state_kernel_caches\n            from .tilelang_decode import clear_tilelang_runtime_caches\n\n            clear_tilelang_kernel_caches()\n            clear_tilelang_state_kernel_caches()\n            clear_tilelang_runtime_caches()\n\n    def prepare_inference_weights(self) -> dict[str, int]:\n        """Install shared TileLang layouts without duplicating model parameters.\n\n        The canonical parameters become views of the read-only inference layouts.\n        ``clear_inference_caches`` restores independent contiguous parameters before\n        training, device moves, state-dict serialization, or backend changes.\n        """\n        if self.kernel_backend != "tilelang":\n            raise RuntimeError("inference layouts require explicit TileLang backend")\n        if self.training:\n            raise RuntimeError("inference layouts require eval mode")\n        reference = self.emb.weight\n        if (\n            reference.device.type != "cuda"\n            or reference.dtype not in {FLOAT16, BFLOAT16}\n            or torch.cuda.get_device_capability(reference.device) != (12, 0)\n        ):\n            raise RuntimeError(\n                "inference layouts require FP16/BF16 parameters on SM120"\n            )\n        created = False\n        with torch.inference_mode():\n            sample = reference.new_empty((1, self.config.hidden_size))\n            for block in cast(Any, self.blocks):\n                *_, attention_created = block.att._prepare_tilelang_cache(sample)\n                *_, ffn_created = block.ffn._tilelang_cache(sample)\n                created = created or attention_created or ffn_created\n        if created:\n            self._inference_cache_epoch += 1\n        return self.inference_cache_stats()\n\n    def state_dict(self, *args: Any, **kwargs: Any) -> dict[str, Any]:\n        if hasattr(self, "_inference_cache_epoch"):\n            self.clear_inference_caches()\n        return cast(dict[str, Any], super().state_dict(*args, **kwargs))\n\n    def load_state_dict(self, *args: Any, **kwargs: Any) -> Any:\n        if hasattr(self, "_inference_cache_epoch"):\n            self.clear_inference_caches()\n        return super().load_state_dict(*args, **kwargs)\n\n    def set_kernel_backend(self, backend: str) -> None:\n        if backend not in {"auto", "torch", "tilelang"}:\n            raise ValueError(f"Unsupported kernel backend: {backend}")\n        if backend != self.kernel_backend:\n            self.clear_inference_caches()\n        self.kernel_backend = backend\n        self.config.kernel_backend = backend\n        for block in self.blocks:\n            if isinstance(block, RWKV7Block):\n                block.att.kernel_backend = backend\n                block.ffn.kernel_backend = backend\n\n    def train(self, mode: bool = True) -> "RWKV7ForCausalLM":\n        if mode and hasattr(self, "_inference_cache_epoch"):\n            self.clear_inference_caches()\n        return cast("RWKV7ForCausalLM", super().train(mode))\n\n    def _apply(self, fn: Any, recurse: bool = True) -> "RWKV7ForCausalLM":\n        if hasattr(self, "_inference_cache_epoch"):\n            self.clear_inference_caches()\n        return cast("RWKV7ForCausalLM", super()._apply(fn, recurse=recurse))\n\n\n    def set_cuda_graph_time_mix(self, enabled: bool) -> None:\n        """Enable exact warm CUDA-graph replay for one-token time-mix inputs."""\n        if enabled:\n            dynamo_config = __import__(\n                "torch._dynamo.config", fromlist=["cache_size_limit"]\n            )\n            minimum_limit = max(64, len(self.blocks) * 2)\n            for name in ("cache_size_limit", "recompile_limit"):\n                current = getattr(dynamo_config, name, 0)\n                if current < minimum_limit:\n                    setattr(dynamo_config, name, minimum_limit)\n        for block in cast(Any, self.blocks):\n            block.att.set_cuda_graph_time_mix(enabled)\n\n\n    def init_state(self, batch_size: int, device, dtype) -> RWKV7State:\n        matrix_dtype = {\n            "float32": FLOAT32,\n            "float16": FLOAT16,\n            "bfloat16": BFLOAT16,\n        }[self.config.recurrent_state_dtype]\n        return RWKV7State.empty(\n            num_layers=self.config.num_hidden_layers,\n            batch_size=batch_size,\n            hidden_size=self.config.hidden_size,\n            num_heads=self.config.num_attention_heads,\n            head_size=self.config.head_size,\n            device=device,\n            dtype=dtype,\n            matrix_dtype=matrix_dtype,\n        )\n\n    def prepare_inputs_for_generation(\n        self,\n        input_ids: torch.Tensor,\n        next_sequence_length: int | None = None,\n        past_key_values: Cache | None = None,\n        attention_mask: torch.Tensor | None = None,\n        inputs_embeds: torch.Tensor | None = None,\n        cache_position: torch.Tensor | None = None,\n        is_first_iteration: bool | None = False,\n        **kwargs: Any,\n    ) -> dict[str, Any]:\n        del next_sequence_length, cache_position, is_first_iteration\n        if past_key_values is not None:\n            input_ids = input_ids[:, -1:]\n            inputs_embeds = None\n        model_inputs: dict[str, Any] = {\n            "input_ids": input_ids,\n            "past_key_values": past_key_values,\n            "attention_mask": attention_mask,\n            "use_cache": kwargs.pop("use_cache", True),\n            **kwargs,\n        }\n        if inputs_embeds is not None and past_key_values is None:\n            model_inputs["inputs_embeds"] = inputs_embeds\n            model_inputs["input_ids"] = None\n        return model_inputs\n\n    def _forward_token_sequence(\n        self,\n        embeds: torch.Tensor,\n        state: RWKV7State,\n        sequence_mask: torch.Tensor | None,\n        reset_mask: torch.Tensor | None,\n    ) -> tuple[torch.Tensor, RWKV7State]:\n        sequence_length = embeds.shape[1]\n        hidden_steps: list[torch.Tensor] = []\n        for token_index in range(sequence_length):\n            x = embeds[:, token_index]\n            active = None\n            if sequence_mask is not None:\n                active = sequence_mask[:, token_index].unsqueeze(-1)\n            reset = None\n            if reset_mask is not None:\n                reset = reset_mask[:, token_index].unsqueeze(-1)\n            _reset_elapsed_tokens(state.elapsed_tokens, reset)\n            first_block = self.blocks[0]\n            if not isinstance(first_block, RWKV7Block):\n                raise TypeError("Invalid first RWKV block")\n            x = first_block.ln0(x)\n            first_value = x\n            next_layers: list[RWKV7LayerState] = []\n            for block, layer_state in zip(\n                self.blocks, state.layer_states, strict=True\n            ):\n                if not isinstance(block, RWKV7Block):\n                    raise TypeError("Invalid RWKV block")\n                layer_state = _reset_layer_state(layer_state, reset)\n                if self.gradient_checkpointing and self.training:\n                    checkpoint = self._gradient_checkpointing_func\n\n                    def block_step(\n                        hidden: torch.Tensor,\n                        channel: torch.Tensor,\n                        time_shift: torch.Tensor,\n                        time_matrix: torch.Tensor,\n                        first: torch.Tensor,\n                        *,\n                        current_block: RWKV7Block = block,\n                    ) -> tuple[\n                        torch.Tensor,\n                        torch.Tensor,\n                        torch.Tensor,\n                        torch.Tensor,\n                        torch.Tensor,\n                    ]:\n                        next_hidden, current_state, next_first = current_block(\n                            hidden,\n                            RWKV7LayerState(channel, time_shift, time_matrix),\n                            first,\n                            active,\n                            state.elapsed_tokens,\n                        )\n                        return (\n                            next_hidden,\n                            current_state.channel,\n                            current_state.time_shift,\n                            current_state.time_matrix,\n                            next_first,\n                        )\n\n                    x, channel, time_shift, time_matrix, first_value = checkpoint(\n                        block_step,\n                        x,\n                        layer_state.channel,\n                        layer_state.time_shift,\n                        layer_state.time_matrix,\n                        first_value,\n                    )\n                    next_state = RWKV7LayerState(\n                        channel, time_shift, time_matrix\n                    )\n                else:\n                    x, next_state, first_value = block(\n                        x,\n                        layer_state,\n                        first_value,\n                        active,\n                        state.elapsed_tokens,\n                    )\n                next_layers.append(next_state)\n            _advance_elapsed_tokens(state.elapsed_tokens, active)\n            state = RWKV7State(\n                next_layers,\n                state.seen_tokens + 1,\n                state.elapsed_tokens,\n            )\n            hidden_steps.append(self.ln_out(x))\n        return TORCH_STACK(hidden_steps, dim=1), state\n\n    def _forward_layer_sequence(\n        self,\n        embeds: torch.Tensor,\n        state: RWKV7State,\n        sequence_mask: torch.Tensor | None,\n        reset_mask: torch.Tensor | None,\n        tokenwise_projections: bool,\n        state_scan_backend: str = "torch",\n    ) -> tuple[torch.Tensor, RWKV7State]:\n        first_block = self.blocks[0]\n        if not isinstance(first_block, RWKV7Block):\n            raise TypeError("Invalid first RWKV block")\n        x = first_block.ln0(embeds)\n        first_value = x\n        active = None if sequence_mask is None else sequence_mask.unsqueeze(-1)\n        reset = None if reset_mask is None else reset_mask.unsqueeze(-1)\n        next_layers: list[RWKV7LayerState] = []\n        for block, layer_state in zip(self.blocks, state.layer_states, strict=True):\n            if not isinstance(block, RWKV7Block):\n                raise TypeError("Invalid RWKV block")\n            if self.gradient_checkpointing and self.training:\n                checkpoint = self._gradient_checkpointing_func\n\n                def block_sequence(\n                    hidden: torch.Tensor,\n                    channel: torch.Tensor,\n                    time_shift: torch.Tensor,\n                    time_matrix: torch.Tensor,\n                    first: torch.Tensor,\n                    *,\n                    current_block: RWKV7Block = block,\n                ) -> tuple[\n                    torch.Tensor,\n                    torch.Tensor,\n                    torch.Tensor,\n                    torch.Tensor,\n                    torch.Tensor,\n                ]:\n                    next_hidden, current_state, next_first = (\n                        current_block.forward_sequence(\n                            hidden,\n                            RWKV7LayerState(channel, time_shift, time_matrix),\n                            first,\n                            active,\n                            reset,\n                            tokenwise_projections,\n                            state_scan_backend,\n                            state.elapsed_tokens,\n                        )\n                    )\n                    return (\n                        next_hidden,\n                        current_state.channel,\n                        current_state.time_shift,\n                        current_state.time_matrix,\n                        next_first,\n                    )\n\n                x, channel, time_shift, time_matrix, first_value = checkpoint(\n                    block_sequence,\n                    x,\n                    layer_state.channel,\n                    layer_state.time_shift,\n                    layer_state.time_matrix,\n                    first_value,\n                )\n                next_state = RWKV7LayerState(channel, time_shift, time_matrix)\n            else:\n                x, next_state, first_value = block.forward_sequence(\n                    x,\n                    layer_state,\n                    first_value,\n                    active,\n                    reset,\n                    tokenwise_projections,\n                    state_scan_backend,\n                    state.elapsed_tokens,\n                )\n            next_layers.append(next_state)\n        elapsed_tokens = state.elapsed_tokens\n        for token_index in range(embeds.shape[1]):\n            token_active = None if active is None else active[:, token_index]\n            token_reset = None if reset is None else reset[:, token_index]\n            _reset_elapsed_tokens(elapsed_tokens, token_reset)\n            _advance_elapsed_tokens(elapsed_tokens, token_active)\n        next_state = RWKV7State(\n            next_layers,\n            state.seen_tokens + embeds.shape[1],\n            elapsed_tokens,\n        )\n        return self.ln_out(x), next_state\n\n    def _forward_chunked_layer_sequence(\n        self,\n        embeds: torch.Tensor,\n        state: RWKV7State,\n        sequence_mask: torch.Tensor | None,\n        reset_mask: torch.Tensor | None,\n        *,\n        chunk_size: int,\n        tokenwise_projections: bool,\n        state_scan_backend: str,\n        hidden_to_keep: int | None,\n    ) -> tuple[torch.Tensor, RWKV7State]:\n        hidden_chunks: list[torch.Tensor] = []\n        retained_hidden: torch.Tensor | None = None\n        for start in range(0, embeds.shape[1], chunk_size):\n            stop = min(start + chunk_size, embeds.shape[1])\n            chunk_mask = (\n                None if sequence_mask is None else sequence_mask[:, start:stop]\n            )\n            chunk_reset = (\n                None if reset_mask is None else reset_mask[:, start:stop]\n            )\n            chunk_hidden, state = self._forward_layer_sequence(\n                embeds[:, start:stop],\n                state,\n                chunk_mask,\n                chunk_reset,\n                tokenwise_projections=tokenwise_projections,\n                state_scan_backend=state_scan_backend,\n            )\n            if hidden_to_keep is None:\n                hidden_chunks.append(chunk_hidden)\n            elif retained_hidden is None:\n                retained_hidden = chunk_hidden[:, -hidden_to_keep:]\n            else:\n                retained_hidden = TORCH_CAT(\n                    (retained_hidden, chunk_hidden), dim=1\n                )[:, -hidden_to_keep:]\n        if hidden_to_keep is not None:\n            if retained_hidden is None:\n                raise RuntimeError("Chunked prefill produced no hidden states")\n            return retained_hidden, state\n        return TORCH_CAT(hidden_chunks, dim=1), state\n\n    def forward(\n        self,\n        input_ids: torch.LongTensor | None = None,\n        attention_mask: torch.Tensor | None = None,\n        labels: torch.LongTensor | None = None,\n        past_key_values: RWKV7State | None = None,\n        use_cache: bool | None = None,\n        inputs_embeds: torch.Tensor | None = None,\n        return_dict: bool | None = None,\n        logits_to_keep: int | torch.Tensor = 0,\n        state_reset_mask: torch.Tensor | None = None,\n        sequence_mode: str = "auto",\n        rwkv_prefill_chunk_size: int | None = None,\n        **kwargs: Any,\n    ) -> Any:\n        if "prefill_chunk_size" in kwargs:\n            raise TypeError(\n                "Use rwkv_prefill_chunk_size; prefill_chunk_size is reserved by "\n                "Transformers generation"\n            )\n        del kwargs\n        if (input_ids is None) == (inputs_embeds is None):\n            raise ValueError("Pass exactly one of input_ids or inputs_embeds")\n        if inputs_embeds is None:\n            if input_ids is None:\n                raise ValueError("input_ids is required when inputs_embeds is absent")\n            embeds = self.emb(input_ids)\n        else:\n            embeds = inputs_embeds\n        if embeds.ndim != 3:\n            raise ValueError(\n                f"Expected [batch, sequence, hidden], got {tuple(embeds.shape)}"\n            )\n\n        batch_size, sequence_length, _ = embeds.shape\n        if sequence_length == 0:\n            raise ValueError("RWKV requires at least one input token")\n        if past_key_values is None:\n            state = self.init_state(batch_size, embeds.device, embeds.dtype)\n        elif not isinstance(past_key_values, RWKV7State):\n            raise TypeError("past_key_values must be RWKV7State")\n        else:\n            state = past_key_values\n        if state.max_batch_size != batch_size:\n            raise ValueError("RWKV state batch size does not match input batch")\n\n        if attention_mask is not None:\n            attention_mask = attention_mask[:, -sequence_length:].to(dtype=BOOL)\n        sequence_mask = attention_mask\n        if state_reset_mask is not None:\n            state_reset_mask = state_reset_mask[:, -sequence_length:].to(dtype=BOOL)\n        if sequence_mode not in {\n            "auto",\n            "token",\n            "layer",\n            "layer-exact",\n            "tilelang-scan",\n            "tilelang-scan-fast",\n        }:\n            raise ValueError(\n                "sequence_mode must be auto, token, layer, layer-exact, "\n                "tilelang-scan or tilelang-scan-fast"\n            )\n        use_tilelang_prefill = (\n            sequence_mode == "auto"\n            and sequence_length > 1\n            and not self.training\n            and not IS_GRAD_ENABLED()\n            and sequence_mask is None\n            and state_reset_mask is None\n            and batch_size == 1\n            and embeds.dtype == FLOAT16\n            and embeds.device.type == "cuda"\n            and torch.cuda.get_device_capability(embeds.device) == (12, 0)\n            and state.elapsed_tokens is not None\n            and state.elapsed_tokens.dtype == torch.int32\n            and state.elapsed_tokens.is_contiguous()\n            and all(\n                layer.time_matrix.dtype == FLOAT16\n                and layer.time_matrix.is_contiguous()\n                for layer in state.layer_states\n            )\n            and self.kernel_backend == "tilelang"\n        )\n        if sequence_mode == "auto":\n            resolved_mode = (\n                "token"\n                if sequence_length == 1\n                else "tilelang-prefill"\n                if use_tilelang_prefill\n                else "layer-exact"\n            )\n        else:\n            resolved_mode = sequence_mode\n        if resolved_mode in {\n            "tilelang-scan",\n            "tilelang-scan-fast",\n        } and self.training:\n            raise RuntimeError("TileLang sequence scan is inference-only")\n        tokenwise_projections = resolved_mode in {\n            "layer-exact",\n            "tilelang-scan",\n            "tilelang-scan-fast",\n        }\n        state_scan_backend = (\n            "tilelang-wkv"\n            if resolved_mode == "tilelang-prefill"\n            else "tilelang-fast"\n            if resolved_mode == "tilelang-scan-fast"\n            else "tilelang"\n            if resolved_mode == "tilelang-scan"\n            else "torch"\n        )\n\n        requested_chunk_size = (\n            0\n            if rwkv_prefill_chunk_size is None and self.training\n            else self.config.rwkv_prefill_chunk_size\n            if rwkv_prefill_chunk_size is None\n            else rwkv_prefill_chunk_size\n        )\n        if requested_chunk_size < 0:\n            raise ValueError("rwkv_prefill_chunk_size must be non-negative")\n        hidden_to_keep = None\n        if (\n            labels is None\n            and isinstance(logits_to_keep, int)\n            and logits_to_keep > 0\n        ):\n            hidden_to_keep = min(logits_to_keep, sequence_length)\n\n        if resolved_mode == "token":\n            hidden, state = self._forward_token_sequence(\n                embeds, state, sequence_mask, state_reset_mask\n            )\n        elif 0 < requested_chunk_size < sequence_length:\n            hidden, state = self._forward_chunked_layer_sequence(\n                embeds,\n                state,\n                sequence_mask,\n                state_reset_mask,\n                chunk_size=requested_chunk_size,\n                tokenwise_projections=tokenwise_projections,\n                state_scan_backend=state_scan_backend,\n                hidden_to_keep=hidden_to_keep,\n            )\n        else:\n            hidden, state = self._forward_layer_sequence(\n                embeds,\n                state,\n                sequence_mask,\n                state_reset_mask,\n                tokenwise_projections=tokenwise_projections,\n                state_scan_backend=state_scan_backend,\n            )\n\n        loss: Any = None\n        head_input = hidden\n        if labels is None:\n            if isinstance(logits_to_keep, int) and logits_to_keep > 0:\n                head_input = hidden[:, -logits_to_keep:]\n            elif isinstance(logits_to_keep, torch.Tensor):\n                head_input = hidden.index_select(\n                    1, logits_to_keep.to(device=hidden.device)\n                )\n        logits = self.head(head_input)\n        if labels is not None:\n            if labels.shape != logits.shape[:2]:\n                raise ValueError("labels must match input_ids shape")\n            effective_labels = labels\n            if sequence_mask is not None:\n                effective_labels = labels.masked_fill(sequence_mask == 0, -100)\n            if state_reset_mask is not None:\n                effective_labels = effective_labels.masked_fill(\n                    state_reset_mask, -100\n                )\n            shift_logits = logits[:, :-1].contiguous().float()\n            shift_labels = effective_labels[:, 1:].contiguous()\n            loss = F.cross_entropy(\n                shift_logits.view(-1, self.config.vocab_size),\n                shift_labels.view(-1),\n                ignore_index=-100,\n            )\n        use_cache = self.config.use_cache if use_cache is None else use_cache\n        output_state = state if use_cache else None\n        return_dict = (\n            bool(getattr(self.config, "return_dict", True))\n            if return_dict is None\n            else return_dict\n        )\n        if not return_dict:\n            output = (logits, output_state)\n            return ((loss,) + output) if loss is not None else output\n        return CausalLMOutputWithPast(\n            loss=loss,\n            logits=logits,\n            past_key_values=output_state,\n        )\n', 'state': 'from __future__ import annotations\n\nfrom contextlib import contextmanager\nfrom dataclasses import dataclass\nfrom threading import Lock\nfrom typing import Callable, Iterator, cast\nimport torch\nfrom transformers.cache_utils import Cache\n\nTORCH_ZEROS = torch.zeros  # type: ignore[attr-defined]\nFLOAT32 = torch.float32  # type: ignore[attr-defined]\nINT32 = torch.int32  # type: ignore[attr-defined]\n\n\n@dataclass\nclass RWKV7LayerState:\n    channel: torch.Tensor\n    time_shift: torch.Tensor\n    time_matrix: torch.Tensor\n\n    def detach(self) -> "RWKV7LayerState":\n        return RWKV7LayerState(\n            self.channel.detach(),\n            self.time_shift.detach(),\n            self.time_matrix.detach(),\n        )\n\n\nclass RWKV7State(Cache):\n    """Explicit recurrent state used as Hugging Face generation cache."""\n\n    def __init__(\n        self,\n        layer_states: list[RWKV7LayerState],\n        seen_tokens: int = 0,\n        elapsed_tokens: torch.Tensor | None = None,\n    ) -> None:\n        super().__init__(layers=[])\n        self.layer_states = layer_states\n        self.seen_tokens = seen_tokens\n        if elapsed_tokens is None and layer_states:\n            reference = layer_states[0].channel\n            elapsed_tokens = torch.full(\n                (reference.shape[0],),\n                seen_tokens,\n                device=reference.device,\n                dtype=INT32,\n            )\n        self.elapsed_tokens = elapsed_tokens\n\n    @classmethod\n    def empty(\n        cls,\n        *,\n        num_layers: int,\n        batch_size: int,\n        hidden_size: int,\n        num_heads: int,\n        head_size: int,\n        device,\n        dtype,\n        matrix_dtype=None,\n    ) -> "RWKV7State":\n        layers = [\n            RWKV7LayerState(\n                channel=TORCH_ZEROS(\n                    batch_size, hidden_size, device=device, dtype=dtype\n                ),\n                time_shift=TORCH_ZEROS(\n                    batch_size, hidden_size, device=device, dtype=dtype\n                ),\n                time_matrix=TORCH_ZEROS(\n                    batch_size,\n                    num_heads,\n                    head_size,\n                    head_size,\n                    device=device,\n                    dtype=matrix_dtype or FLOAT32,\n                ),\n            )\n            for _ in range(num_layers)\n        ]\n        return cls(\n            layers,\n            elapsed_tokens=TORCH_ZEROS(\n                batch_size, device=device, dtype=INT32\n            ),\n        )\n\n    @property\n    def is_compileable(self) -> bool:\n        return False\n\n    @property\n    def is_initialized(self) -> bool:\n        return bool(self.layer_states)\n\n    @property\n    def max_batch_size(self) -> int:\n        if not self.layer_states:\n            return 0\n        return self.layer_states[0].channel.shape[0]\n\n    @property\n    def max_cache_len(self) -> int:\n        return self.seen_tokens\n\n    def get_seq_length(self, layer_idx: int = 0) -> int:\n        del layer_idx\n        return self.seen_tokens\n\n    def get_max_cache_shape(self, layer_idx: int = 0) -> int:\n        del layer_idx\n        return self.seen_tokens\n\n    def detach(self) -> "RWKV7State":\n        return RWKV7State(\n            [layer.detach() for layer in self.layer_states],\n            self.seen_tokens,\n            None if self.elapsed_tokens is None else self.elapsed_tokens.detach(),\n        )\n\n    def clone(self) -> "RWKV7State":\n        return RWKV7State(\n            [\n                RWKV7LayerState(\n                    layer.channel.clone(),\n                    layer.time_shift.clone(),\n                    layer.time_matrix.clone(),\n                )\n                for layer in self.layer_states\n            ],\n            self.seen_tokens,\n            None if self.elapsed_tokens is None else self.elapsed_tokens.clone(),\n        )\n\n    def copy_(self, source: "RWKV7State") -> "RWKV7State":\n        """Copy another same-layout state without changing destination addresses."""\n        if len(self.layer_states) != len(source.layer_states):\n            raise ValueError("RWKV states have different layer counts")\n        for destination, current in zip(\n            self.layer_states, source.layer_states, strict=True\n):\n            for target, value in (\n                (destination.channel, current.channel),\n                (destination.time_shift, current.time_shift),\n                (destination.time_matrix, current.time_matrix),\n            ):\n                if (\n                    target.shape != value.shape\n                    or target.dtype != value.dtype\n                    or target.device != value.device\n                ):\n                    raise ValueError("RWKV states have different tensor layouts")\n                target.copy_(value)\n        if (self.elapsed_tokens is None) != (source.elapsed_tokens is None):\n            raise ValueError("RWKV states disagree on elapsed-token storage")\n        if self.elapsed_tokens is not None and source.elapsed_tokens is not None:\n            if (\n                self.elapsed_tokens.shape != source.elapsed_tokens.shape\n                or self.elapsed_tokens.device != source.elapsed_tokens.device\n            ):\n                raise ValueError("RWKV elapsed-token layouts differ")\n            self.elapsed_tokens.copy_(source.elapsed_tokens)\n        self.seen_tokens = source.seen_tokens\n        return self\n\n    def copy_batch_row_(\n        self, source: "RWKV7State", index: int, *, seen_tokens: int\n) -> "RWKV7State":\n        """Copy one row from a batched state into this stable one-request state."""\n        if self.max_batch_size != 1:\n            raise ValueError("destination state must contain exactly one request")\n        if index < 0 or index >= source.max_batch_size:\n            raise IndexError("source state row is out of range")\n        if len(self.layer_states) != len(source.layer_states):\n            raise ValueError("RWKV states have different layer counts")\n        for destination, current in zip(\n            self.layer_states, source.layer_states, strict=True\n):\n            for target, value in (\n                (destination.channel, current.channel[index : index + 1]),\n                (destination.time_shift, current.time_shift[index : index + 1]),\n                (destination.time_matrix, current.time_matrix[index : index + 1]),\n            ):\n                if (\n                    target.shape != value.shape\n                    or target.dtype != value.dtype\n                    or target.device != value.device\n                ):\n                    raise ValueError("RWKV states have different tensor layouts")\n                target.copy_(value)\n        if (self.elapsed_tokens is None) != (source.elapsed_tokens is None):\n            raise ValueError("RWKV states disagree on elapsed-token storage")\n        if self.elapsed_tokens is not None and source.elapsed_tokens is not None:\n            self.elapsed_tokens.copy_(source.elapsed_tokens[index : index + 1])\n        self.seen_tokens = seen_tokens\n        return self\n\n    @classmethod\n    def batch_stack(cls, states: list["RWKV7State"]) -> "RWKV7State":\n        """Stack independent request states into one decode batch."""\n        if not states:\n            raise ValueError("at least one RWKV state is required")\n        layer_count = len(states[0].layer_states)\n        elapsed_present = states[0].elapsed_tokens is not None\n        if any(len(state.layer_states) != layer_count for state in states):\n            raise ValueError("RWKV states have different layer counts")\n        if any(\n            (state.elapsed_tokens is not None) != elapsed_present for state in states\n):\n            raise ValueError("RWKV states disagree on elapsed-token storage")\n        layers: list[RWKV7LayerState] = []\n        for layer_index in range(layer_count):\n            current = [state.layer_states[layer_index] for state in states]\n            reference = current[0]\n            for layer in current[1:]:\n                for value, expected in (\n                    (layer.channel, reference.channel),\n                    (layer.time_shift, reference.time_shift),\n                    (layer.time_matrix, reference.time_matrix),\n                ):\n                    if (\n                        value.shape[1:] != expected.shape[1:]\n                        or value.dtype != expected.dtype\n                        or value.device != expected.device\n                    ):\n                        raise ValueError("RWKV states have incompatible tensor layouts")\n            layers.append(\n                RWKV7LayerState(\n                    torch.cat([layer.channel for layer in current], dim=0),\n                    torch.cat([layer.time_shift for layer in current], dim=0),\n                    torch.cat([layer.time_matrix for layer in current], dim=0),\n                )\n            )\n        elapsed = (\n            torch.cat(\n                [cast(torch.Tensor, state.elapsed_tokens) for state in states],\n                dim=0,\n            )\n            if elapsed_present\n            else None\n        )\n        return cls(layers, max(state.seen_tokens for state in states), elapsed)\n\n    def batch_split(\n        self,\n        batch_sizes: list[int] | None = None,\n        *,\n        seen_tokens: list[int] | None = None,\n    ) -> list["RWKV7State"]:\n        """Split a batched state into cloned independent request states."""\n        if batch_sizes is None:\n            batch_sizes = [1] * self.max_batch_size\n        if not batch_sizes or any(size < 1 for size in batch_sizes):\n            raise ValueError("batch_sizes must contain positive values")\n        if sum(batch_sizes) != self.max_batch_size:\n            raise ValueError("batch_sizes do not cover the RWKV state batch")\n        if seen_tokens is not None and len(seen_tokens) != len(batch_sizes):\n            raise ValueError("seen_tokens must match batch_sizes")\n        outputs: list[RWKV7State] = []\n        start = 0\n        for index, size in enumerate(batch_sizes):\n            stop = start + size\n            outputs.append(\n                RWKV7State(\n                    [\n                        RWKV7LayerState(\n                            layer.channel[start:stop].clone(),\n                            layer.time_shift[start:stop].clone(),\n                            layer.time_matrix[start:stop].clone(),\n                        )\n                        for layer in self.layer_states\n                    ],\n                    self.seen_tokens\n                    if seen_tokens is None\n                    else seen_tokens[index],\n                    None\n                    if self.elapsed_tokens is None\n                    else self.elapsed_tokens[start:stop].clone(),\n                )\n            )\n            start = stop\n        return outputs\n\n    def reorder_cache(self, beam_idx: torch.LongTensor) -> None:\n        self.batch_select_indices(beam_idx)\n\n    def batch_repeat_interleave(self, repeats: int) -> None:\n        for layer in self.layer_states:\n            layer.channel = layer.channel.repeat_interleave(repeats, dim=0)\n            layer.time_shift = layer.time_shift.repeat_interleave(repeats, dim=0)\n            layer.time_matrix = layer.time_matrix.repeat_interleave(repeats, dim=0)\n        if self.elapsed_tokens is not None:\n            self.elapsed_tokens = self.elapsed_tokens.repeat_interleave(repeats, dim=0)\n    def batch_select_indices(self, indices: torch.Tensor) -> None:\n        for layer in self.layer_states:\n            device_indices = indices.to(layer.channel.device)\n            layer.channel = layer.channel.index_select(0, device_indices)\n            layer.time_shift = layer.time_shift.index_select(0, device_indices)\n            layer.time_matrix = layer.time_matrix.index_select(0, device_indices)\n        if self.elapsed_tokens is not None:\n            device_indices = indices.to(self.elapsed_tokens.device)\n            self.elapsed_tokens = self.elapsed_tokens.index_select(0, device_indices)\n    def crop(self, max_length: int) -> None:\n        if max_length < self.seen_tokens:\n            raise NotImplementedError(\n                "RWKV recurrent state cannot roll back without state history"\n            )\n\n    def reset(self) -> None:\n        for layer in self.layer_states:\n            layer.channel.zero_()\n            layer.time_shift.zero_()\n            layer.time_matrix.zero_()\n        if self.elapsed_tokens is not None:\n            self.elapsed_tokens.zero_()\n        self.seen_tokens = 0\n\n\nclass RWKV7StatePool:\n    """Bounded pool that zeroes recurrent state between requests."""\n\n    def __init__(\n        self, factory: Callable[[], RWKV7State], *, max_entries: int = 8\n    ) -> None:\n        if max_entries < 1:\n            raise ValueError("max_entries must be positive")\n        self.factory = factory\n        self.max_entries = max_entries\n        prototype = factory()\n        with torch.inference_mode():\n            prototype.reset()\n        self._signature = self._state_signature(prototype)\n        self._available = [prototype]\n        self._leased: dict[int, RWKV7State] = {}\n        self._total = 1\n        self._closed = False\n        self._lock = Lock()\n\n    @staticmethod\n    def _state_signature(state: RWKV7State) -> tuple:\n        tensors = [\n            tensor\n            for layer in state.layer_states\n            for tensor in (layer.channel, layer.time_shift, layer.time_matrix)\n        ]\n        if state.elapsed_tokens is not None:\n            tensors.append(state.elapsed_tokens)\n        return (\n            len(state.layer_states),\n            state.elapsed_tokens is not None,\n            tuple(\n                (tuple(tensor.shape), tensor.dtype, str(tensor.device))\n                for tensor in tensors\n            ),\n        )\n\n    def acquire(self) -> RWKV7State:\n        with self._lock:\n            if self._closed:\n                raise RuntimeError("RWKV state pool is closed")\n            if self._available:\n                state = self._available.pop()\n            elif self._total < self.max_entries:\n                state = self.factory()\n                if self._state_signature(state) != self._signature:\n                    raise RuntimeError("RWKV state factory changed tensor layout")\n                self._total += 1\n            else:\n                raise RuntimeError("RWKV state pool is exhausted")\n            with torch.inference_mode():\n                state.reset()\n            self._leased[id(state)] = state\n            return state\n\n    def release(self, state: RWKV7State) -> None:\n        with self._lock:\n            leased = self._leased.pop(id(state), None)\n            if leased is not state:\n                raise ValueError("state is not leased from this pool")\n            if self._state_signature(state) != self._signature:\n                self._total -= 1\n                raise ValueError("state tensor layout changed while leased")\n            with torch.inference_mode():\n                state.reset()\n            if self._closed:\n                self._total -= 1\n            else:\n                self._available.append(state)\n\n    def gather(self, states: list[RWKV7State]) -> RWKV7State:\n        """Gather leased request states into a contiguous decode batch."""\n        with self._lock:\n            if any(self._leased.get(id(state)) is not state for state in states):\n                raise ValueError("all states must be leased from this pool")\n            return type(states[0]).batch_stack(states)\n\n    def scatter_(\n        self,\n        source: RWKV7State,\n        states: list[RWKV7State],\n        *,\n        seen_tokens: list[int],\n    ) -> None:\n        """Scatter a decode batch directly into stable leased destinations."""\n        if len(states) != source.max_batch_size or len(seen_tokens) != len(states):\n            raise ValueError("scatter rows must match leased states")\n        with self._lock, torch.inference_mode():\n            if any(self._leased.get(id(state)) is not state for state in states):\n                raise ValueError("all states must be leased from this pool")\n            for index, (state, seen) in enumerate(zip(states, seen_tokens, strict=True)):\n                state.copy_batch_row_(source, index, seen_tokens=seen)\n    @contextmanager\n    def lease(self) -> Iterator[RWKV7State]:\n        state = self.acquire()\n        try:\n            yield state\n        finally:\n            self.release(state)\n\n    def clear(self) -> None:\n        with self._lock:\n            self._total -= len(self._available)\n            self._available.clear()\n\n    def close(self) -> None:\n        with self._lock:\n            self._closed = True\n            self._total -= len(self._available)\n            self._available.clear()\n\n    def stats(self) -> dict[str, int]:\n        with self._lock:\n            return {\n                "max_entries": self.max_entries,\n                "total": self._total,\n                "available": len(self._available),\n                "leased": len(self._leased),\n            }\n\n\n@dataclass\nclass _RWKV7StatePage:\n    page_id: int\n    state: RWKV7State\n    free_slots: list[int]\n    leased_slots: set[int]\n    allocated_bytes: int\n\n    @property\n    def slots(self) -> int:\n        return self.state.max_batch_size\n\n\nclass RWKV7PagedStatePool:\n    """Bounded lazy page allocator for constant-size RWKV recurrent states."""\n\n    def __init__(\n        self,\n        factory: Callable[[], RWKV7State],\n        *,\n        max_entries: int = 8,\n        page_size: int = 8,\n        contiguous_views: bool = True,\n    ) -> None:\n        if max_entries < 1:\n            raise ValueError("max_entries must be positive")\n        if page_size < 1:\n            raise ValueError("page_size must be positive")\n        prototype = factory()\n        if prototype.max_batch_size != 1:\n            raise ValueError("paged state factory must create one-request states")\n        with torch.inference_mode():\n            prototype.reset()\n        self.factory = factory\n        self.max_entries = max_entries\n        self.page_size = min(page_size, max_entries)\n        self.contiguous_views = contiguous_views\n        self.max_pages = (max_entries + self.page_size - 1) // self.page_size\n        self._signature = RWKV7StatePool._state_signature(prototype)\n        self._state_type = type(prototype)\n        self._layer_type = type(prototype.layer_states[0]) if prototype.layer_states else RWKV7LayerState\n        self._layer_specs = [\n            tuple(\n                (tuple(tensor.shape[1:]), tensor.dtype, tensor.device)\n                for tensor in (layer.channel, layer.time_shift, layer.time_matrix)\n            )\n            for layer in prototype.layer_states\n        ]\n        self._elapsed_spec = (\n            None\n            if prototype.elapsed_tokens is None\n            else (prototype.elapsed_tokens.dtype, prototype.elapsed_tokens.device)\n        )\n        self._pages: dict[int, _RWKV7StatePage] = {}\n        self._leased: dict[int, tuple[RWKV7State, int, int]] = {}\n        self._next_page_id = 0\n        self._closed = False\n        self._contiguous_view_batches = 0\n        self._gather_copy_batches = 0\n        self._gathered_rows = 0\n        self._lock = Lock()\n\n    @staticmethod\n    def _state_bytes(state: RWKV7State) -> int:\n        tensors = [\n            tensor\n            for layer in state.layer_states\n            for tensor in (layer.channel, layer.time_shift, layer.time_matrix)\n        ]\n        if state.elapsed_tokens is not None:\n            tensors.append(state.elapsed_tokens)\n        return sum(tensor.numel() * tensor.element_size() for tensor in tensors)\n\n    def _allocate_page(self) -> _RWKV7StatePage:\n        allocated_slots = sum(page.slots for page in self._pages.values())\n        slots = min(self.page_size, self.max_entries - allocated_slots)\n        if slots < 1:\n            raise RuntimeError("RWKV paged state pool is exhausted")\n        layers = []\n        for specs in self._layer_specs:\n            tensors = [\n                torch.zeros((slots, *shape), dtype=dtype, device=device)\n                for shape, dtype, device in specs\n            ]\n            layers.append(self._layer_type(*tensors))\n        elapsed = (\n            None\n            if self._elapsed_spec is None\n            else torch.zeros(\n                slots, dtype=self._elapsed_spec[0], device=self._elapsed_spec[1]\n            )\n        )\n        page_state = self._state_type(layers, 0, elapsed)\n        page = _RWKV7StatePage(\n            page_id=self._next_page_id,\n            state=page_state,\n            free_slots=list(range(slots - 1, -1, -1)),\n            leased_slots=set(),\n            allocated_bytes=self._state_bytes(page_state),\n        )\n        self._next_page_id += 1\n        self._pages[page.page_id] = page\n        return page\n\n    def _slot_view(self, page: _RWKV7StatePage, slot: int) -> RWKV7State:\n        layers = [\n            self._layer_type(\n                layer.channel[slot : slot + 1],\n                layer.time_shift[slot : slot + 1],\n                layer.time_matrix[slot : slot + 1],\n            )\n            for layer in page.state.layer_states\n        ]\n        elapsed = (\n            None\n            if page.state.elapsed_tokens is None\n            else page.state.elapsed_tokens[slot : slot + 1]\n        )\n        return self._state_type(layers, 0, elapsed)\n\n    def acquire(self) -> RWKV7State:\n        with self._lock:\n            if self._closed:\n                raise RuntimeError("RWKV paged state pool is closed")\n            page = next(\n                (current for current in self._pages.values() if current.free_slots),\n                None,\n            )\n            if page is None:\n                page = self._allocate_page()\n            slot = page.free_slots.pop()\n            page.leased_slots.add(slot)\n            state = self._slot_view(page, slot)\n            with torch.inference_mode():\n                state.reset()\n            self._leased[id(state)] = (state, page.page_id, slot)\n            return state\n\n    def release(self, state: RWKV7State) -> None:\n        with self._lock:\n            lease = self._leased.pop(id(state), None)\n            if lease is None or lease[0] is not state:\n                raise ValueError("state is not leased from this paged pool")\n            _, page_id, slot = lease\n            page = self._pages[page_id]\n            layout_valid = RWKV7StatePool._state_signature(state) == self._signature\n            stored = self._slot_view(page, slot)\n            with torch.inference_mode():\n                stored.reset()\n            page.leased_slots.remove(slot)\n            if self._closed and not page.leased_slots:\n                del self._pages[page_id]\n            else:\n                page.free_slots.append(slot)\n            if not layout_valid:\n                raise ValueError("state tensor layout changed while leased")\n\n    @contextmanager\n    def lease(self) -> Iterator[RWKV7State]:\n        state = self.acquire()\n        try:\n            yield state\n        finally:\n            self.release(state)\n\n    def _contiguous_batch_view_locked(\n        self, states: list[RWKV7State]\n    ) -> RWKV7State | None:\n        leases = [self._leased[id(state)] for state in states]\n        page_id = leases[0][1]\n        if any(lease[1] != page_id for lease in leases):\n            return None\n        slots = [lease[2] for lease in leases]\n        start = slots[0]\n        if slots != list(range(start, start + len(slots))):\n            return None\n        page = self._pages[page_id]\n        stop = start + len(slots)\n        layers = [\n            self._layer_type(\n                layer.channel[start:stop],\n                layer.time_shift[start:stop],\n                layer.time_matrix[start:stop],\n            )\n            for layer in page.state.layer_states\n        ]\n        elapsed = (\n            None\n            if page.state.elapsed_tokens is None\n            else page.state.elapsed_tokens[start:stop]\n        )\n        return self._state_type(\n            layers, max(state.seen_tokens for state in states), elapsed\n        )\n\n    def gather(self, states: list[RWKV7State]) -> RWKV7State:\n        """Return a zero-copy contiguous page view or a copied fallback batch."""\n        if not states:\n            raise ValueError("at least one leased state is required")\n        with self._lock:\n            if any(\n                (lease := self._leased.get(id(state))) is None\n                or lease[0] is not state\n                for state in states\n            ):\n                raise ValueError("all states must be leased from this paged pool")\n            self._gathered_rows += len(states)\n            view = (\n                self._contiguous_batch_view_locked(states)\n                if self.contiguous_views\n                else None\n            )\n            if view is not None:\n                self._contiguous_view_batches += 1\n                return view\n            self._gather_copy_batches += 1\n            return type(states[0]).batch_stack(states)\n\n    def scatter_(\n        self,\n        source: RWKV7State,\n        states: list[RWKV7State],\n        *,\n        seen_tokens: list[int],\n    ) -> None:\n        if len(states) != source.max_batch_size or len(seen_tokens) != len(states):\n            raise ValueError("scatter rows must match leased states")\n        with self._lock, torch.inference_mode():\n            if any(\n                (lease := self._leased.get(id(state))) is None or lease[0] is not state\n                for state in states\n            ):\n                raise ValueError("all states must be leased from this paged pool")\n            for index, (state, seen) in enumerate(zip(states, seen_tokens, strict=True)):\n                state.copy_batch_row_(source, index, seen_tokens=seen)\n\n    def lease_info(self, state: RWKV7State) -> dict[str, int]:\n        with self._lock:\n            lease = self._leased.get(id(state))\n            if lease is None or lease[0] is not state:\n                raise ValueError("state is not leased from this paged pool")\n            return {"page_id": lease[1], "slot": lease[2]}\n\n    def clear(self) -> None:\n        with self._lock:\n            empty_pages = [\n                page_id\n                for page_id, page in self._pages.items()\n                if not page.leased_slots\n            ]\n            for page_id in empty_pages:\n                del self._pages[page_id]\n\n    def close(self) -> None:\n        with self._lock:\n            self._closed = True\n            empty_pages = [\n                page_id\n                for page_id, page in self._pages.items()\n                if not page.leased_slots\n            ]\n            for page_id in empty_pages:\n                del self._pages[page_id]\n\n    def stats(self) -> dict[str, int | str]:\n        with self._lock:\n            return {\n                "kind": "paged",\n                "max_entries": self.max_entries,\n                "page_size": self.page_size,\n                "contiguous_views": int(self.contiguous_views),\n                "max_pages": self.max_pages,\n                "pages": len(self._pages),\n                "total": sum(page.slots for page in self._pages.values()),\n                "available": sum(len(page.free_slots) for page in self._pages.values()),\n                "leased": len(self._leased),\n                "contiguous_view_batches": self._contiguous_view_batches,\n                "gather_copy_batches": self._gather_copy_batches,\n                "gathered_rows": self._gathered_rows,\n                "allocated_bytes": sum(\n                    page.allocated_bytes for page in self._pages.values()\n                ),\n            }\n', 'tilelang_decode': 'from __future__ import annotations\n\nfrom dataclasses import dataclass\nfrom typing import Any, Sequence\n\nimport torch\n\nfrom .kernel import decode as _kernel_decode\n_compiled_cmix_add_layernorm_mix = _kernel_decode._compiled_cmix_add_layernorm_mix\n_compiled_cmix_finalize = _kernel_decode._compiled_cmix_finalize\n_compiled_cmix_layernorm_mix = _kernel_decode._compiled_cmix_layernorm_mix\n_compiled_cmix_value = _kernel_decode._compiled_cmix_value\n_compiled_ffn = _kernel_decode._compiled_ffn\n_compiled_gemv = _kernel_decode._compiled_gemv\n_compiled_cmix_binned_finalize = _kernel_decode._compiled_cmix_binned_finalize\n_compiled_cmix_sparse_binned = _kernel_decode._compiled_cmix_sparse_binned\n_compiled_key_gate = _kernel_decode._compiled_key_gate\n_compiled_post_state = _kernel_decode._compiled_post_state\n_compiled_tmix_layernorm_mix6 = _kernel_decode._compiled_tmix_layernorm_mix6\n_compiled_rankout = _kernel_decode._compiled_rankout\n_compiled_rankout_reduced = _kernel_decode._compiled_rankout_reduced\n_compiled_rkv = _kernel_decode._compiled_rkv\n_wkv_out = _kernel_decode._wkv_out\n_wkv_w0_t1_out = _kernel_decode._wkv_w0_t1_out\nfrom .kernel import state as _kernel_state\ncuda_arch_key = _kernel_state.cuda_arch_key\n\n_STREAM_POOLS: dict[int, tuple[torch.cuda.Stream, ...]] = {}\n\n\ndef _stream_pool(device: torch.device) -> tuple[torch.cuda.Stream, ...]:\n    index = device.index\n    if index is None:\n        index = torch.cuda.current_device()\n    streams = _STREAM_POOLS.get(index)\n    if streams is None:\n        with torch.cuda.device(index):\n            streams = tuple(torch.cuda.Stream() for _ in range(4))\n        _STREAM_POOLS[index] = streams\n    return streams\n\n\n@dataclass(frozen=True)\nclass TileLangDecodeSpec:\n    """Fixed B1T1 dimensions for optimized inference."""\n\n    channels: int\n    ffn_rows: int\n    num_heads: int\n    head_size: int\n    ranks: tuple[int, int, int, int]\n\n    def __post_init__(self) -> None:\n        if self.channels <= 0 or self.ffn_rows <= 0:\n            raise ValueError("decode dimensions must be positive")\n        if self.num_heads * self.head_size != self.channels:\n            raise ValueError("num_heads * head_size must equal channels")\n        if self.head_size != 64:\n            raise ValueError("TileLang WKV requires head size 64")\n        if len(self.ranks) != 4 or any(rank < 0 for rank in self.ranks):\n            raise ValueError("ranks must contain four non-negative values")\n\n    @property\n    def rkv_output_rows(self) -> int:\n        return 3 * self.channels + sum(self.ranks)\n\n\n@dataclass\nclass TileLangDecodeWorkspace:\n    """Caller-owned graph-capturable workspace for one layer/request."""\n\n    rkv_output: torch.Tensor\n    rank_decay: torch.Tensor\n    rank_gate_a: torch.Tensor\n    rank_gate_g: torch.Tensor\n    rank_value: torch.Tensor\n    wkv_output: torch.Tensor\n    attention_output: torch.Tensor\n    normalized_key: torch.Tensor\n    modified_key: torch.Tensor\n    anti_key: torch.Tensor\n    anti_gate: torch.Tensor\n    block_residual: torch.Tensor\n    tmix_normalized: torch.Tensor\n    tmix_mixed: torch.Tensor\n    post_mixed: torch.Tensor\n    ffn_normalized: torch.Tensor\n    ffn_hidden: torch.Tensor\n    ffn_partials: torch.Tensor\n    ffn_output: torch.Tensor\n    ffn_bins: torch.Tensor\n\n\n\nclass TileLangDecodeBackend:\n    """Single ultra-optimized SM120 TileLang B1T1 inference backend.\n\n    PyTorch remains pure reference implementation. Callers own packed weights,\n    recurrent state, and workspaces; backend performs no hidden allocation.\n    """\n\n    def __init__(\n        self,\n        spec: TileLangDecodeSpec,\n        device: torch.device | str,\n        dtype: torch.dtype = torch.float16,\n    ):\n        self.spec = spec\n        requested_device = torch.device(device)\n        self.device = (\n            torch.device("cuda", torch.cuda.current_device())\n            if requested_device.type == "cuda" and requested_device.index is None\n            else requested_device\n        )\n        if self.device.type != "cuda":\n            raise RuntimeError("TileLang decode backend requires CUDA")\n        if torch.cuda.get_device_capability(self.device) != (12, 0):\n            raise RuntimeError("TileLang decode backend requires SM120")\n        if dtype not in {torch.float16, torch.bfloat16}:\n            raise TypeError("TileLang decode backend requires FP16 or BF16")\n        self.dtype = dtype\n        self.input_dtype = "float16" if dtype == torch.float16 else "bfloat16"\n        self._architecture = cuda_arch_key(self.device)\n        self._rkv_kernel: Any | None = None\n        self._gemv_kernels: dict[tuple[int, int, bool], Any] = {}\n        self._ffn_kernel: Any | None = None\n        self._cmix_sparse_binned_kernel: Any | None = None\n        self._cmix_binned_finalize_kernel: Any | None = None\n        self._cmix_layernorm_mix_kernel: Any | None = None\n        self._cmix_add_layernorm_mix_kernel: Any | None = None\n        self._tmix_layernorm_mix6_kernel: Any | None = None\n        self._key_gate_kernel: Any | None = None\n        self._cmix_value_kernel: Any | None = None\n        self._cmix_finalize_kernel: Any | None = None\n        self._rkv_binding: tuple[int, ...] | None = None\n        self._ffn_binding: tuple[int, ...] | None = None\n        self._rankout_kernels: dict[bool, Any] = {}\n        self._rankout_bindings: dict[bool, tuple[int, ...]] = {}\n        self._rankout_reduced_kernels: dict[bool, Any] = {}\n        self._rankout_reduced_bindings: dict[bool, tuple[int, ...]] = {}\n        self._post_state_kernel: Any | None = None\n        self._post_state_binding: tuple[int, ...] | None = None\n        self._streams = _stream_pool(self.device)\n\n    def create_workspace(self) -> TileLangDecodeWorkspace:\n        options = {"device": self.device, "dtype": self.dtype}\n        return TileLangDecodeWorkspace(\n            rkv_output=torch.empty(self.spec.rkv_output_rows, **options),\n            rank_decay=torch.empty(self.spec.channels, **options),\n            rank_gate_a=torch.empty(self.spec.channels, **options),\n            rank_gate_g=torch.empty(self.spec.channels, **options),\n            rank_value=torch.empty(self.spec.channels, **options),\n            wkv_output=torch.empty(\n                (1, 1, self.spec.num_heads, self.spec.head_size), **options\n            ),\n            attention_output=torch.empty(self.spec.channels, **options),\n            normalized_key=torch.empty(self.spec.channels, **options),\n            modified_key=torch.empty(self.spec.channels, **options),\n            anti_key=torch.empty(self.spec.channels, **options),\n            anti_gate=torch.empty(self.spec.channels, **options),\n            block_residual=torch.empty(self.spec.channels, **options),\n            tmix_normalized=torch.empty(self.spec.channels, **options),\n            tmix_mixed=torch.empty((6, self.spec.channels), **options),\n            post_mixed=torch.empty(self.spec.channels, **options),\n            ffn_normalized=torch.empty(self.spec.channels, **options),\n            ffn_hidden=torch.empty(self.spec.ffn_rows, **options),\n            ffn_partials=torch.empty(\n                (4, self.spec.channels), **options\n            ),\n            ffn_output=torch.empty(self.spec.channels, **options),\n            ffn_bins=torch.zeros(\n                (6, self.spec.channels),\n                device=self.device,\n                dtype=torch.float32,\n            ),\n        )\n\n    def pack_rkv_weights(\n        self,\n        rkv_weights: Sequence[torch.Tensor],\n        lowrank_weights: Sequence[torch.Tensor],\n    ) -> torch.Tensor:\n        """Pack three dense and four low-rank weights once during model load."""\n        if len(rkv_weights) != 3 or len(lowrank_weights) != 4:\n            raise ValueError("expected three RKV and four low-rank weights")\n        expected_rows = (self.spec.channels,) * 3 + self.spec.ranks\n        weights = tuple(rkv_weights) + tuple(lowrank_weights)\n        for weight, rows in zip(weights, expected_rows, strict=True):\n            self._validate(weight, (rows, self.spec.channels), "weight")\n        return torch.cat(weights, dim=0).contiguous()\n\n    def pack_ffn_value_weight(self, weight: torch.Tensor) -> torch.Tensor:\n        """Pack FFN value weight as activation-major contiguous tiles."""\n        self._validate(\n            weight,\n            (self.spec.channels, self.spec.ffn_rows),\n            "FFN value weight",\n        )\n        return weight.t().contiguous()\n\n    def rkv(\n        self,\n        inputs: Sequence[torch.Tensor],\n        packed_weight: torch.Tensor,\n        output: torch.Tensor,\n    ) -> None:\n        """Run direct-input R/K/V and W/A/G/V rank-input projection."""\n        if len(inputs) != 6:\n            raise ValueError("expected xr, xk, xv, xw, xa, and xg")\n        tensors = (*inputs, packed_weight, output)\n        self._reject_training(tensors)\n        if self._rkv_kernel is None:\n            self._rkv_kernel = _compiled_rkv(\n                self.spec.channels,\n                *self.spec.ranks,\n                2,\n                128,\n                self.input_dtype,\n                self._architecture,\n            )\n        binding = tuple(tensor.data_ptr() for tensor in tensors)\n        if binding != self._rkv_binding:\n            for value in inputs:\n                self._validate(value, (self.spec.channels,), "RKV input")\n            self._validate(\n                packed_weight,\n                (self.spec.rkv_output_rows, self.spec.channels),\n                "packed RKV weight",\n            )\n            self._validate(\n                output, (self.spec.rkv_output_rows,), "RKV output"\n            )\n            self._rkv_binding = binding\n        self._rkv_kernel(*inputs, packed_weight, output)\n\n    def gemv(\n        self,\n        value: torch.Tensor,\n        weight: torch.Tensor,\n        output: torch.Tensor,\n        clear_output: torch.Tensor | None = None,\n    ) -> None:\n        """Run output-tiled inference GEMV into caller-owned output."""\n        input_rows = value.numel()\n        output_rows = output.numel()\n        clear_target = value if clear_output is None else clear_output\n        tensors = (value, weight, output, clear_target)\n        self._reject_training(tensors)\n        key = (input_rows, output_rows, clear_output is not None)\n        kernel = self._gemv_kernels.get(key)\n        if kernel is None:\n            kernel = _compiled_gemv(\n                input_rows,\n                output_rows,\n                self.input_dtype,\n                2,\n                128,\n                clear_output is not None,\n                self._architecture,\n            )\n            self._gemv_kernels[key] = kernel\n        self._validate(value, (input_rows,), "GEMV input")\n        self._validate(weight, (output_rows, input_rows), "GEMV weight")\n        self._validate(output, (output_rows,), "GEMV output")\n        self._validate(clear_target, (input_rows,), "GEMV clear output")\n        kernel(*tensors)\n\n    def rkv_views(self, output: torch.Tensor) -> tuple[torch.Tensor, ...]:\n        """Return zero-copy R, K, V, W1, A1, G1, V1 views."""\n        sizes = (self.spec.channels,) * 3 + self.spec.ranks\n        return tuple(output.split(sizes))\n\n    def rankout(\n        self,\n        ranks: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],\n        weights: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],\n        vectors: tuple[torch.Tensor, torch.Tensor, torch.Tensor],\n        value_base: torch.Tensor,\n        first_value: torch.Tensor,\n        workspace: TileLangDecodeWorkspace,\n        *,\n        use_value_mix: bool,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n        """Fuse W/A/G/V rank-out and pointwise finalization."""\n        outputs = (\n            workspace.rank_decay,\n            workspace.rank_gate_a,\n            workspace.rank_gate_g,\n            workspace.rank_value,\n        )\n        tensors = (*ranks, *weights, *vectors, value_base, first_value, *outputs)\n        self._reject_training(tensors)\n        kernel = self._rankout_kernels.get(use_value_mix)\n        if kernel is None:\n            kernel = _compiled_rankout(\n                self.spec.channels,\n                *self.spec.ranks,\n                use_value_mix,\n                self.input_dtype,\n                self._architecture,\n            )\n            self._rankout_kernels[use_value_mix] = kernel\n        binding = tuple(tensor.data_ptr() for tensor in tensors)\n        if self._rankout_bindings.get(use_value_mix) != binding:\n            for rank, size in zip(ranks, self.spec.ranks, strict=True):\n                self._validate(rank, (size,), "rank-out input")\n            for weight, size in zip(weights, self.spec.ranks, strict=True):\n                self._validate(\n                    weight, (size, self.spec.channels), "rank-out weight"\n                )\n            for vector in vectors:\n                self._validate(\n                    vector, (self.spec.channels,), "rank-out vector"\n                )\n            self._validate(\n                value_base, (self.spec.channels,), "rank-out value"\n            )\n            self._validate(\n                first_value, (self.spec.channels,), "rank-out first value"\n            )\n            self._rankout_bindings[use_value_mix] = binding\n        kernel(*ranks, *weights, *vectors, value_base, first_value, *outputs)\n        return outputs\n\n    def rankout_reduced(\n        self,\n        ranks: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],\n        weights: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],\n        vectors: tuple[torch.Tensor, torch.Tensor],\n        value_base: torch.Tensor,\n        first_value: torch.Tensor,\n        workspace: TileLangDecodeWorkspace,\n        *,\n        use_value_mix: bool,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n        """Run rank-parallel W/A/G/V output from transposed weights."""\n        outputs = (\n            workspace.rank_decay,\n            workspace.rank_gate_a,\n            workspace.rank_gate_g,\n            workspace.rank_value,\n        )\n        tensors = (*ranks, *weights, *vectors, value_base, first_value, *outputs)\n        self._reject_training(tensors)\n        kernel = self._rankout_reduced_kernels.get(use_value_mix)\n        if kernel is None:\n            kernel = _compiled_rankout_reduced(\n                self.spec.channels,\n                *self.spec.ranks,\n                use_value_mix,\n                self.input_dtype,\n                4,\n                128,\n                self._architecture,\n            )\n            self._rankout_reduced_kernels[use_value_mix] = kernel\n        binding = tuple(tensor.data_ptr() for tensor in tensors)\n        if self._rankout_reduced_bindings.get(use_value_mix) != binding:\n            for rank, size in zip(ranks, self.spec.ranks, strict=True):\n                self._validate(rank, (size,), "reduced rank-out input")\n            for weight, size in zip(weights, self.spec.ranks, strict=True):\n                self._validate(\n                    weight,\n                    (self.spec.channels, size),\n                    "transposed reduced rank-out weight",\n                )\n            for vector in vectors:\n                self._validate(\n                    vector, (self.spec.channels,), "reduced rank-out vector"\n                )\n            self._validate(\n                value_base, (self.spec.channels,), "reduced rank-out value"\n            )\n            self._validate(\n                first_value, (self.spec.channels,), "reduced rank-out first value"\n            )\n            self._rankout_reduced_bindings[use_value_mix] = binding\n        kernel(*ranks, *weights, *vectors, value_base, first_value, *outputs)\n        return outputs\n\n\n    def key_gate(\n        self,\n        key: torch.Tensor,\n        key_scale: torch.Tensor,\n        gate_a: torch.Tensor,\n        gate_scale: torch.Tensor,\n        workspace: TileLangDecodeWorkspace,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n        """Fuse per-head key normalization and recurrent gate vectors."""\n        outputs = (\n            workspace.normalized_key,\n            workspace.modified_key,\n            workspace.anti_key,\n            workspace.anti_gate,\n        )\n        tensors = (key, key_scale, gate_a, gate_scale, *outputs)\n        self._reject_training(tensors)\n        if self._key_gate_kernel is None:\n            self._key_gate_kernel = _compiled_key_gate(\n                self.spec.num_heads,\n                self.spec.head_size,\n                self.input_dtype,\n                self._architecture,\n            )\n        for tensor in tensors:\n            self._validate(\n                tensor, (self.spec.channels,), "key normalization/gate tensor"\n            )\n        self._key_gate_kernel(*tensors)\n        return outputs\n\n    def tmix_layernorm_mix6(\n        self,\n        residual: torch.Tensor,\n        previous: torch.Tensor,\n        norm_weight: torch.Tensor,\n        norm_bias: torch.Tensor,\n        mix_weights: torch.Tensor,\n        epsilon: float,\n        workspace: TileLangDecodeWorkspace,\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        """Fuse LayerNorm and six shifted time-mix vectors."""\n        normalized = previous\n        mixed = workspace.tmix_mixed\n        tensors = (\n            residual,\n            previous,\n            norm_weight,\n            norm_bias,\n            mix_weights,\n            normalized,\n            mixed,\n        )\n        self._reject_training(tensors)\n        if self._tmix_layernorm_mix6_kernel is None:\n            self._tmix_layernorm_mix6_kernel = _compiled_tmix_layernorm_mix6(\n                self.spec.channels,\n                self.input_dtype,\n                epsilon,\n                256,\n                self._architecture,\n            )\n        for tensor in tensors[:-1]:\n            expected = (\n                (6, self.spec.channels)\n                if tensor is mix_weights\n                else (self.spec.channels,)\n            )\n            self._validate(tensor, expected, "time-mix LayerNorm tensor")\n        self._validate(\n            mixed, (6, self.spec.channels), "time-mix mixed output"\n        )\n        self._tmix_layernorm_mix6_kernel(*tensors)\n        return normalized, mixed\n\n    def post_state(\n        self,\n        projected: torch.Tensor,\n        receptance: torch.Tensor,\n        key: torch.Tensor,\n        value: torch.Tensor,\n        r_k: torch.Tensor,\n        gate: torch.Tensor,\n        norm_weight: torch.Tensor,\n        norm_bias: torch.Tensor,\n        epsilon: float,\n        workspace: TileLangDecodeWorkspace,\n    ) -> torch.Tensor:\n        """Fuse GroupNorm, RKV residual, and gate finalization."""\n        output = workspace.post_mixed\n        tensors = (\n            projected,\n            receptance,\n            key,\n            value,\n            r_k,\n            gate,\n            norm_weight,\n            norm_bias,\n            output,\n        )\n        self._reject_training(tensors)\n        if self._post_state_kernel is None:\n            self._post_state_kernel = _compiled_post_state(\n                self.spec.num_heads,\n                self.spec.head_size,\n                self.input_dtype,\n                epsilon,\n                self._architecture,\n            )\n        binding = tuple(tensor.data_ptr() for tensor in tensors)\n        if self._post_state_binding != binding:\n            head_shape = (self.spec.num_heads, self.spec.head_size)\n            for tensor in (projected, receptance, key, value, r_k):\n                self._validate(tensor, head_shape, "post-state head tensor")\n            for tensor in (gate, norm_weight, norm_bias, output):\n                self._validate(\n                    tensor, (self.spec.channels,), "post-state channel tensor"\n                )\n            self._post_state_binding = binding\n        self._post_state_kernel(*tensors)\n        return output\n\n\n    def cmix_layernorm_mix(\n        self,\n        residual: torch.Tensor,\n        previous: torch.Tensor,\n        norm_weight: torch.Tensor,\n        norm_bias: torch.Tensor,\n        mix_weight: torch.Tensor,\n        epsilon: float,\n        workspace: TileLangDecodeWorkspace,\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        """Fuse LayerNorm and shifted channel-mix input generation."""\n        normalized = workspace.ffn_normalized\n        mixed = workspace.post_mixed\n        tensors = (\n            residual,\n            previous,\n            norm_weight,\n            norm_bias,\n            mix_weight,\n            normalized,\n            mixed,\n        )\n        self._reject_training(tensors)\n        if self._cmix_layernorm_mix_kernel is None:\n            self._cmix_layernorm_mix_kernel = _compiled_cmix_layernorm_mix(\n                self.spec.channels,\n                self.input_dtype,\n                epsilon,\n                256,\n                self._architecture,\n            )\n        for tensor in tensors:\n            self._validate(\n                tensor, (self.spec.channels,), "channel-mix LayerNorm tensor"\n            )\n        self._cmix_layernorm_mix_kernel(*tensors)\n        return normalized, mixed\n\n    def cmix_add_layernorm_mix(\n        self,\n        residual: torch.Tensor,\n        update: torch.Tensor,\n        previous: torch.Tensor,\n        norm_weight: torch.Tensor,\n        norm_bias: torch.Tensor,\n        mix_weight: torch.Tensor,\n        epsilon: float,\n        workspace: TileLangDecodeWorkspace,\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n        """Fuse residual add, LayerNorm, and shifted channel-mix input."""\n        combined = workspace.block_residual\n        normalized = previous\n        mixed = workspace.post_mixed\n        tensors = (\n            residual,\n            update,\n            previous,\n            norm_weight,\n            norm_bias,\n            mix_weight,\n            combined,\n            normalized,\n            mixed,\n        )\n        self._reject_training(tensors)\n        if self._cmix_add_layernorm_mix_kernel is None:\n            self._cmix_add_layernorm_mix_kernel = (\n                _compiled_cmix_add_layernorm_mix(\n                    self.spec.channels,\n                    self.input_dtype,\n                    epsilon,\n                    256,\n                    self._architecture,\n                )\n            )\n        for tensor in tensors:\n            self._validate(\n                tensor, (self.spec.channels,), "channel-mix add/LayerNorm tensor"\n            )\n        self._cmix_add_layernorm_mix_kernel(*tensors)\n        return combined, normalized, mixed\n\n    def ffn(\n        self,\n        mixed: torch.Tensor,\n        key_weight: torch.Tensor,\n        packed_value_weight: torch.Tensor,\n        output: torch.Tensor,\n        workspace: TileLangDecodeWorkspace,\n        residual: torch.Tensor | None = None,\n    ) -> None:\n        """Run tiled key GEMV plus sparse TileLang down projection."""\n        tensors = (mixed, key_weight, packed_value_weight, output)\n        if residual is not None:\n            tensors = (*tensors, residual)\n        self._reject_training(tensors)\n        self._validate(mixed, (self.spec.channels,), "FFN mixed")\n        self._validate(\n            key_weight,\n            (self.spec.ffn_rows, self.spec.channels),\n            "FFN key weight",\n        )\n        expected_value_shape = (self.spec.ffn_rows, self.spec.channels)\n        self._validate(\n            packed_value_weight,\n            expected_value_shape,\n            "FFN value weight",\n        )\n        self._validate(output, (self.spec.channels,), "FFN output")\n        if residual is not None:\n            self._validate(residual, (self.spec.channels,), "FFN residual")\n\n        if self.dtype == torch.float16:\n            if self._cmix_sparse_binned_kernel is None:\n                self._cmix_sparse_binned_kernel = _compiled_cmix_sparse_binned(\n                    self.spec.channels,\n                    self.spec.ffn_rows,\n                    self.input_dtype,\n                    128,\n                    128,\n                    6,\n                    self._architecture,\n                )\n                self._cmix_binned_finalize_kernel = (\n                    _compiled_cmix_binned_finalize(\n                        self.spec.channels,\n                        6,\n                        self.input_dtype,\n                        256,\n                        self._architecture,\n                    )\n                )\n            self.gemv(mixed, key_weight, workspace.ffn_hidden)\n            if self._cmix_binned_finalize_kernel is None:\n                raise RuntimeError("binned FFN finalizer failed to initialize")\n            self._cmix_sparse_binned_kernel(\n                workspace.ffn_hidden,\n                packed_value_weight,\n                workspace.ffn_bins,\n            )\n            residual_input = output if residual is None else residual\n            self._cmix_binned_finalize_kernel(\n                workspace.ffn_bins, residual_input, output\n            )\n            return\n\n        split_rows = self.spec.ffn_rows // 4\n        if self.spec.ffn_rows % 4:\n            raise ValueError("FFN rows must be divisible by four streams")\n        if self._ffn_kernel is None:\n            self._ffn_kernel = _compiled_ffn(\n                self.spec.channels,\n                split_rows,\n                2,\n                128,\n                self.input_dtype,\n                self._architecture,\n            )\n            self._cmix_value_kernel = _compiled_cmix_value(\n                self.spec.channels,\n                split_rows,\n                self.input_dtype,\n                64,\n                8,\n                self._architecture,\n            )\n            self._cmix_finalize_kernel = _compiled_cmix_finalize(\n                self.spec.channels,\n                4,\n                self.input_dtype,\n                self._architecture,\n            )\n        if (\n            self._ffn_kernel is None\n            or self._cmix_value_kernel is None\n            or self._cmix_finalize_kernel is None\n        ):\n            raise RuntimeError("TileLang BF16 FFN kernels failed to initialize")\n        current_stream = torch.cuda.current_stream(self.device)\n        for split, stream in enumerate(self._streams):\n            start = split * split_rows\n            stop = start + split_rows\n            stream.wait_stream(current_stream)\n            with torch.cuda.stream(stream):\n                self._ffn_kernel(\n                    mixed,\n                    key_weight[start:stop],\n                    workspace.ffn_hidden[start:stop],\n                )\n                self._cmix_value_kernel(\n                    workspace.ffn_hidden[start:stop],\n                    packed_value_weight[start:stop],\n                    workspace.ffn_partials[split],\n                )\n        for stream in self._streams:\n            current_stream.wait_stream(stream)\n        self._cmix_finalize_kernel(workspace.ffn_partials, output)\n\n    def wkv_w0_t1(\n        self,\n        state: torch.Tensor,\n        receptance: torch.Tensor,\n        decay: torch.Tensor,\n        decay_bias: torch.Tensor,\n        key: torch.Tensor,\n        value: torch.Tensor,\n        gate_a: torch.Tensor,\n        gate_b: torch.Tensor,\n        elapsed: torch.Tensor,\n        output: torch.Tensor,\n    ) -> None:\n        """Run specialized B1T1 FP16 WKV with fused static decay bias."""\n        self._reject_training(\n            (\n                state,\n                receptance,\n                decay,\n                decay_bias,\n                key,\n                value,\n                gate_a,\n                gate_b,\n                output,\n            )\n        )\n        _wkv_w0_t1_out(\n            state,\n            receptance,\n            decay,\n            decay_bias,\n            key,\n            value,\n            gate_a,\n            gate_b,\n            elapsed,\n            output,\n        )\n\n    def wkv(\n        self,\n        state: torch.Tensor,\n        receptance: torch.Tensor,\n        decay_raw: torch.Tensor,\n        key: torch.Tensor,\n        value: torch.Tensor,\n        gate_a: torch.Tensor,\n        gate_b: torch.Tensor,\n        elapsed: torch.Tensor,\n        output: torch.Tensor,\n    ) -> None:\n        """Run exact FP16 recurrent update into caller-owned buffers."""\n        self._reject_training(\n            (state, receptance, decay_raw, key, value, gate_a, gate_b, output)\n        )\n        _wkv_out(\n            state,\n            receptance,\n            decay_raw,\n            key,\n            value,\n            gate_a,\n            gate_b,\n            elapsed,\n            output,\n        )\n\n    def _validate(\n        self, tensor: torch.Tensor, shape: tuple[int, ...], name: str\n) -> None:\n        if tensor.device != self.device:\n            raise ValueError(f"{name} must be on {self.device}")\n        if tensor.dtype != self.dtype:\n            raise TypeError(f"{name} must use {self.dtype}")\n        if tuple(tensor.shape) != shape:\n            raise ValueError(f"{name} must have shape {shape}")\n        if not tensor.is_contiguous():\n            raise ValueError(f"{name} must be contiguous")\n\n    @staticmethod\n    def _reject_training(tensors: Sequence[torch.Tensor]) -> None:\n        if torch.is_grad_enabled() and any(\n            tensor.requires_grad for tensor in tensors\n        ):\n            raise RuntimeError("TileLang decode backend is inference-only")\n\n\ndef clear_tilelang_runtime_caches() -> None:\n    """Synchronize and release the bounded per-device auxiliary stream pools."""\n    for device_index in tuple(_STREAM_POOLS):\n        with torch.cuda.device(device_index):\n            torch.cuda.synchronize(device_index)\n    _STREAM_POOLS.clear()\n'}


def _load_module(name):
    source = _SOURCES[name]
    module_name = f'{__package__}.{name}'
    filename = f'<rwkv7_runtime_{name}_{SOURCE_SHA256[name]}>'
    linecache.cache[filename] = (
        len(source), None, source.splitlines(keepends=True), filename
    )
    module = ModuleType(module_name)
    module.__file__ = filename
    module.__package__ = __package__
    sys.modules[module_name] = module
    try:
        exec(compile(source, filename, 'exec'), module.__dict__, module.__dict__)
    except Exception:
        sys.modules.pop(module_name, None)
        raise
    return module


_MODULE_ORDER = ('configuration_rwkv7', 'state', 'custom_ops', 'kernel_dispatch', 'tilelang_decode', 'modeling_rwkv7')
_MODULES = {name: _load_module(name) for name in _MODULE_ORDER}
RWKV7Config = _MODULES['configuration_rwkv7'].RWKV7Config
RWKV7LayerState = _MODULES['state'].RWKV7LayerState
RWKV7State = _MODULES['state'].RWKV7State
RWKV7ForCausalLM = _MODULES['modeling_rwkv7'].RWKV7ForCausalLM

__all__ = ["RWKV7Config", "RWKV7ForCausalLM", "RWKV7LayerState", "RWKV7State"]