# ============================================================================= # COPYRIGHT © 2025-2026 Konstantin Vladimirovich Grabko. ALL RIGHTS RESERVED. # CMS Manhattan JiRack Technology — PATENT PENDING # # This code is proprietary. # Personal and non-commercial research use is allowed. # Any commercial use, derivative works for profit, or distribution # requires a paid license and 5% royalty. # # Unauthorized commercial use is strictly prohibited. # Contact: grabko@cmsmanhattan.com # ============================================================================= # inference with KV-Cache support import torch import torch.nn as nn import torch.nn.functional as F # Импортируем базовые классы и функции из твоих прошлых файлов from JiRackTernaryPyTorch_1b import apply_rotary_emb, repeat_kv from JiRackTernaryPyTorch_1b_inf import TernaryTransformer1BInf, TransformerBlockInference class TransformerBlockInferenceKV(TransformerBlockInference): """ A subclass of the output block that overrides the forward method to ensure proper concatenation and preserve key-value changes. """ def __init__(self, config): super().__init__(config) def forward(self, x, freqs_cos, freqs_sin, past_kv=None, position_ids=None): h = self.norm1(x) B, T, D = x.shape # Высчитываем Q, K, V через твои BitLinearInference слои q = self.q_proj(h).view(B, T, self.n_heads, self.head_dim).transpose(1, 2) k = self.k_proj(h).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(h).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) # Применяем RoPE к текущим токенам. # Если передан position_ids (индексация в ONNX), берем срез частот по нему. if position_ids is not None: f_cos = freqs_cos[position_ids].to(xq.device).unsqueeze(1).repeat(1, 1, 1, 2) f_sin = freqs_sin[position_ids].to(xq.device).unsqueeze(1).repeat(1, 1, 1, 2) q = (q * f_cos) + (self._rotate_half(q) * f_sin) k = (k * f_cos) + (self._rotate_half(k) * f_sin) else: q, k = apply_rotary_emb(q, k, freqs_cos, freqs_sin) # --- НАСТОЯЩАЯ СКЛЕЙКА КЭША --- if past_kv is not None: past_k, past_v = past_kv k = torch.cat([past_k, k], dim=2) # Склеиваем по оси seq_len v = torch.cat([past_v, v], dim=2) present_kv = (k, v) # Текущее состояние для возврата наверх # Дублируем KV для Grouped-Query Attention (GQA) k_rep = repeat_kv(k, self.n_rep) v_rep = repeat_kv(v, self.n_rep) # Считаем внимание по всей накопленной истории токенов. # Каузальная маска активна только при первоначальном заполнении (T > 1) attn_out = F.scaled_dot_product_attention(q, k_rep, v_rep, is_causal=(T > 1 and past_kv is None)) x = x + self.out_proj(attn_out.transpose(1, 2).reshape(B, T, D)) m = self.norm2(x) x = x + self.ffn_w2(F.silu(self.ffn_w1(m)) * self.ffn_w3(m)) return x, present_kv def _rotate_half(self, x): x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) class TernaryTransformer1BInfKV(TernaryTransformer1BInf): """ The main inference model class JiRack with KV-cache support. Inherits the weight-loading logic `load_prod_weights` from TernaryTransformer1BInf """ def __init__(self, config): super().__init__(config) # Заменяем стандартные блоки на блоки с поддержкой KV-кэша self.blocks = nn.ModuleList([ TransformerBlockInferenceKV(config) for _ in range(config.num_hidden_layers) ]) def forward(self, input_ids, past_key_values=None, position_ids=None): """ Input parameters: - input_ids: Current tokens [Batch, SeqLen] - past_key_values: List of tuples (past_k, past_v) for each layer - position_ids: Optional position indices for RoPE (required for ONNX) """ x = self.token_emb(input_ids) present_key_values = [] for i, block in enumerate(self.blocks): # Извлекаем кэш конкретного слоя, если он передан block_past = past_key_values[i] if past_key_values is not None else None # Прогоняем через блок x, present_kv = block( x, self.freqs_cos, self.freqs_sin, past_kv=block_past, position_ids=position_ids ) present_key_values.append(present_kv) logits = self.lm_head(self.ln_f(x)) # Возвращаем логиты и обновленный кэш return logits, present_key_values