diff --git "a/inference/runtime.py" "b/inference/runtime.py" new file mode 100644--- /dev/null +++ "b/inference/runtime.py" @@ -0,0 +1,38 @@ +"""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'' + 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"]