""" Context Window — Janela deslizante com cache KV para sequências longas. Implementa três estratégias (ver docs/MATH_ANALYSIS.md seção 3.4): - Sliding Window Pura - Sink + Sliding (StreamingLLM) — RECOMENDADO - Attention Recomputation (sem cache, máxima precisão) Para CNN-BiGRU cooperativo, o contexto é gerenciado em modo: - Encoder (treino/classificação): processa janela inteira, BiGRU bidirecional - Decoder (geração): cache KV para CausalSelfAttention + GRU unidirecional Autor: CNN-BiGRU Project """ from __future__ import annotations import logging from dataclasses import dataclass, field from typing import Dict, List, Optional, Tuple import torch import torch.nn as nn logger = logging.getLogger(__name__) # ============================================================================ # Configuração # ============================================================================ @dataclass class ContextWindowConfig: """Configuração da janela de contexto.""" # Comprimento máximo da janela (em tokens) max_window: int = 512 # Estratégia de evicção: "sliding" | "sink_sliding" | "recompute" eviction_strategy: str = "sink_sliding" # Número de tokens "sink" (apenas para sink_sliding) — tipicamente BOS + system prompt n_sink_tokens: int = 4 # Dimensão do embedding (para alocar cache KV) embed_dim: int = 256 # Número de cabeças (para cache KV em multi-head attention) n_heads: int = 4 # Dimensão por cabeça head_dim: Optional[int] = None # default: embed_dim // n_heads # Número de camadas (para cache KV multicamada) n_layers: int = 2 # Device padrão device: str = "cpu" # Dtype do cache (None = manter fp32) dtype: Optional[torch.dtype] = None # ============================================================================ # KV Cache # ============================================================================ class KVCache: """ Cache de pares (Key, Value) para CausalSelfAttention. Shape por camada: K: [n_heads, seq_cached, head_dim] V: [n_heads, seq_cached, head_dim] Em modo batch: K: [batch, n_heads, seq_cached, head_dim] V: [batch, n_heads, seq_cached, head_dim] """ def __init__(self, n_layers: int, batch_size: int, n_heads: int, head_dim: int, device: torch.device, dtype: torch.dtype): self.n_layers = n_layers self.batch_size = batch_size self.n_heads = n_heads self.head_dim = head_dim self.device = device self.dtype = dtype # Pre-alocar listas vazias; preencher lazy no primeiro update self.keys: List[Optional[torch.Tensor]] = [None] * n_layers self.values: List[Optional[torch.Tensor]] = [None] * n_layers def update(self, layer_idx: int, new_k: torch.Tensor, new_v: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: """ Atualiza o cache da camada layer_idx com novos K, V. Args: layer_idx: índice da camada [0, n_layers) new_k: [batch, n_heads, seq_new, head_dim] new_v: [batch, n_heads, seq_new, head_dim] Returns: (cached_k, cached_v): [batch, n_heads, seq_total, head_dim] """ if not (0 <= layer_idx < self.n_layers): raise IndexError(f"layer_idx {layer_idx} fora de range [0, {self.n_layers})") # Validar shapes if new_k.dim() != 4 or new_v.dim() != 4: raise ValueError( f"new_k e new_v devem ser 4D [batch, n_heads, seq, head_dim]; " f"recebido {new_k.shape}" ) # Converter dtype/device new_k = new_k.to(device=self.device, dtype=self.dtype) new_v = new_v.to(device=self.device, dtype=self.dtype) if self.keys[layer_idx] is None: self.keys[layer_idx] = new_k self.values[layer_idx] = new_v else: self.keys[layer_idx] = torch.cat([self.keys[layer_idx], new_k], dim=2) self.values[layer_idx] = torch.cat([self.values[layer_idx], new_v], dim=2) return self.keys[layer_idx], self.values[layer_idx] def evict_sliding(self, max_keep: int) -> None: """ Estratégia sliding window pura: mantém apenas os últimos max_keep tokens. """ for layer_idx in range(self.n_layers): if self.keys[layer_idx] is None: continue seq = self.keys[layer_idx].size(2) if seq > max_keep: start = seq - max_keep self.keys[layer_idx] = self.keys[layer_idx][:, :, start:, :].contiguous() self.values[layer_idx] = self.values[layer_idx][:, :, start:, :].contiguous() def evict_sink_sliding(self, max_keep: int, n_sink: int) -> None: """ Estratégia sink + sliding: mantém os primeiros n_sink + últimos (max_keep - n_sink). """ if n_sink >= max_keep: logger.warning( f"n_sink ({n_sink}) >= max_keep ({max_keep}); usando sliding puro" ) self.evict_sliding(max_keep) return n_sliding = max_keep - n_sink for layer_idx in range(self.n_layers): if self.keys[layer_idx] is None: continue seq = self.keys[layer_idx].size(2) if seq > max_keep: sink_k = self.keys[layer_idx][:, :, :n_sink, :] sink_v = self.values[layer_idx][:, :, :n_sink, :] sliding_k = self.keys[layer_idx][:, :, -n_sliding:, :] sliding_v = self.values[layer_idx][:, :, -n_sliding:, :] self.keys[layer_idx] = torch.cat([sink_k, sliding_k], dim=2).contiguous() self.values[layer_idx] = torch.cat([sink_v, sliding_v], dim=2).contiguous() def reset(self) -> None: """Limpa o cache completamente.""" self.keys = [None] * self.n_layers self.values = [None] * self.n_layers def get_seq_len(self, layer_idx: int = 0) -> int: if self.keys[layer_idx] is None: return 0 return self.keys[layer_idx].size(2) def total_tokens(self) -> int: return max(self.get_seq_len(i) for i in range(self.n_layers)) # ============================================================================ # Context Window Manager # ============================================================================ class ContextWindowManager: """ Gerencia a janela de contexto aplicando a política de evicção configurada. """ def __init__(self, config: ContextWindowConfig): self.config = config self.head_dim = config.head_dim or (config.embed_dim // config.n_heads) if config.embed_dim % config.n_heads != 0: raise ValueError( f"embed_dim ({config.embed_dim}) deve ser divisível por " f"n_heads ({config.n_heads})" ) self.kv_cache: Optional[KVCache] = None # Sequência de tokens brutos (para recomputação se necessário) self.token_history: List[torch.Tensor] = [] # ---------------------------------------------------------------------- # Inicialização # ---------------------------------------------------------------------- def init_cache(self, batch_size: int, device: torch.device, dtype: Optional[torch.dtype] = None) -> KVCache: """Cria novo cache KV para uma sessão de geração.""" dtype = dtype or self.config.dtype or torch.float32 self.kv_cache = KVCache( n_layers=self.config.n_layers, batch_size=batch_size, n_heads=self.config.n_heads, head_dim=self.head_dim, device=device, dtype=dtype, ) self.token_history = [] return self.kv_cache # ---------------------------------------------------------------------- # Adicionar tokens # ---------------------------------------------------------------------- def append_tokens(self, token_ids: torch.Tensor) -> torch.Tensor: """ Adiciona tokens ao histórico e retorna a janela ativa. Args: token_ids: [batch, seq_new] Returns: window: [batch, seq_window] — tokens na janela ativa após evicção """ if token_ids.dim() != 2: raise ValueError(f"token_ids deve ser [batch, seq]; recebido {token_ids.shape}") # Para simplificação, mantemos histórico apenas do batch[0] # (geração autoregressiva tipicamente batch=1) if token_ids.size(0) > 1: logger.warning("ContextWindowManager.append_tokens: batch>1, histórico rastreia apenas batch[0]") self.token_history.append(token_ids) # Concatenar todo o histórico full = torch.cat(self.token_history, dim=1) # Aplicar evicção se necessário return self._evict_tokens(full) def _evict_tokens(self, full_seq: torch.Tensor) -> torch.Tensor: """Aplica política de evicção à sequência completa.""" seq_len = full_seq.size(1) max_keep = self.config.max_window if seq_len <= max_keep: return full_seq strategy = self.config.eviction_strategy if strategy == "sliding": return full_seq[:, -max_keep:] elif strategy == "sink_sliding": n_sink = self.config.n_sink_tokens n_sliding = max_keep - n_sink sink = full_seq[:, :n_sink] sliding = full_seq[:, -n_sliding:] return torch.cat([sink, sliding], dim=1) elif strategy == "recompute": # Manter tudo (atenção recomputada) return full_seq else: logger.warning(f"Estratégia de evicção desconhecida: {strategy}; usando sliding") return full_seq[:, -max_keep:] # ---------------------------------------------------------------------- # Evict cache KV # ---------------------------------------------------------------------- def evict_cache(self) -> None: """Aplica a política de evicção ao cache KV.""" if self.kv_cache is None: return strategy = self.config.eviction_strategy max_keep = self.config.max_window if strategy == "sliding": self.kv_cache.evict_sliding(max_keep) elif strategy == "sink_sliding": self.kv_cache.evict_sink_sliding(max_keep, self.config.n_sink_tokens) elif strategy == "recompute": self.kv_cache.reset() else: self.kv_cache.evict_sliding(max_keep) # ---------------------------------------------------------------------- # Reset # ---------------------------------------------------------------------- def reset(self) -> None: """Reseta todo o estado (cache + histórico).""" if self.kv_cache is not None: self.kv_cache.reset() self.kv_cache = None self.token_history = [] # ---------------------------------------------------------------------- # Info # ---------------------------------------------------------------------- def get_info(self) -> Dict: """Retorna informações sobre o estado atual.""" cache_len = self.kv_cache.total_tokens() if self.kv_cache else 0 history_len = sum(t.size(1) for t in self.token_history) return { "max_window": self.config.max_window, "strategy": self.config.eviction_strategy, "n_sink": self.config.n_sink_tokens, "cache_tokens": cache_len, "history_tokens": history_len, "active": self.kv_cache is not None, } # ============================================================================ # Helpers # ============================================================================ def make_context_window( max_window: int = 512, strategy: str = "sink_sliding", n_sink: int = 4, embed_dim: int = 256, n_heads: int = 4, n_layers: int = 2, device: str = "cpu", ) -> ContextWindowManager: """Factory para criar ContextWindowManager com defaults sensatos.""" config = ContextWindowConfig( max_window=max_window, eviction_strategy=strategy, n_sink_tokens=n_sink, embed_dim=embed_dim, n_heads=n_heads, head_dim=embed_dim // n_heads, n_layers=n_layers, device=device, ) return ContextWindowManager(config) # ============================================================================ # Contexto de 1M Tokens — Chunked / Ring Attention # ============================================================================ @dataclass class LongContextConfig: """Configuração para contexto de até 1M tokens. Estratégias suportadas: - "chunked": divide a sequência em chunks de tamanho chunk_size, processa cada chunk separadamente, mantém um cache KV global. - "ring": Ring Attention — distribui chunks entre dispositivos/GPUs (apenas para multi-GPU; em CPU faz fallback para chunked). - "sink_sliding_long": sink + sliding com window grande (até max_window). - "hybrid": usa chunked para encoder, sink_sliding para decoder. Para 1M tokens: - Memória necessária para cache KV FP32: n_layers * n_heads * head_dim * 1e6 * 2 (K+V) * 4 bytes = n_layers * d_model * 1e6 * 8 bytes Ex: 6 layers * 256 d * 1e6 * 8 = ~12 GB (inviável em CPU) - Com chunked attention (chunk_size=8192), a memória por chunk é: n_layers * d_model * 8192 * 8 = ~16 MB (viável) - Ring Attention distribui entre N dispositivos: memória / N """ max_window: int = 1_000_000 # 1M tokens chunk_size: int = 8192 # tamanho do chunk para processamento strategy: str = "chunked" # "chunked" | "ring" | "sink_sliding_long" | "hybrid" n_sink_tokens: int = 16 # tokens sink (BOS + system) n_sliding_tokens: int = 8192 # sliding window para sink_sliding_long embed_dim: int = 256 n_heads: int = 4 head_dim: Optional[int] = None n_layers: int = 2 device: str = "cpu" dtype: Optional[torch.dtype] = None # Ring attention n_ring_devices: int = 1 # 1 = CPU single, >1 = multi-GPU # Overlap communication (apenas para ring) overlap_comm: bool = True # Cache para chunks processados (evita recomputação) use_chunk_cache: bool = True chunk_cache_max: int = 128 # máximo de chunks em cache class LongContextManager: """Gerencia contexto de até 1M tokens via chunked / ring attention. Para 1M tokens, a estratégia padrão é: 1. Dividir a sequência em chunks de `chunk_size` tokens 2. Processar cada chunk com atenção local (intra-chunk) 3. Comunicar informações entre chunks via: a. Tokens de "boundary" (compartilhados entre chunks adjacentes) b. Sumarização hierárquica (chunk summary -> global context) 4. Manter cache KV apenas para o chunk atual + sink tokens Em modo "ring" (multi-GPU), cada device processa um chunk e os resultados são propagados em anel (P2P communication). Para CNN-BiGRU: - O BiGRU bidirecional não funciona bem em modo chunked puro (perde dependências backward entre chunks). Solução: processar cada chunk em ambas as direções, comunicar estado hidden entre chunks. - Em modo decoder (geração), usa sink + sliding puro (cache KV). """ def __init__(self, config: LongContextConfig): self.config = config self.head_dim = config.head_dim or (config.embed_dim // config.n_heads) if config.embed_dim % config.n_heads != 0: raise ValueError( f"embed_dim ({config.embed_dim}) deve ser divisível por " f"n_heads ({config.n_heads})" ) self.kv_cache: Optional[KVCache] = None self.token_history: List[torch.Tensor] = [] # Chunk cache: cache de sumários de chunks processados self.chunk_summaries: List[torch.Tensor] = [] # Posições dos tokens para positional encoding self._global_position_offset = 0 def init_cache( self, batch_size: int, device: torch.device, dtype: Optional[torch.dtype] = None, ) -> KVCache: """Cria novo cache KV para sessão de geração longa.""" dtype = dtype or self.config.dtype or torch.float32 self.kv_cache = KVCache( n_layers=self.config.n_layers, batch_size=batch_size, n_heads=self.config.n_heads, head_dim=self.head_dim, device=device, dtype=dtype, ) self.token_history = [] self.chunk_summaries = [] self._global_position_offset = 0 return self.kv_cache def append_tokens_chunked( self, token_ids: torch.Tensor, ) -> Dict[str, torch.Tensor]: """Adiciona tokens usando estratégia chunked. Args: token_ids: [batch, seq_new] Returns: dict com: active_window: [batch, seq_active] tokens ativos chunk_index: int — índice do chunk atual is_new_chunk: bool — True se iniciou novo chunk n_chunks_total: int — total de chunks processados """ if token_ids.dim() != 2: raise ValueError(f"token_ids deve ser [batch, seq]; recebido {token_ids.shape}") self.token_history.append(token_ids) full = torch.cat(self.token_history, dim=1) total_len = full.size(1) # Verificar se ultrapassou chunk boundary chunk_size = self.config.chunk_size n_chunks = (total_len + chunk_size - 1) // chunk_size is_new_chunk = len(self.chunk_summaries) < n_chunks # Janela ativa: último chunk + sink tokens n_sink = self.config.n_sink_tokens n_sliding = self.config.n_sliding_tokens or chunk_size # Sempre manter sink + últimos n_sliding tokens if total_len > n_sink + n_sliding: sink = full[:, :n_sink] sliding = full[:, -n_sliding:] active = torch.cat([sink, sliding], dim=1) else: active = full return { "active_window": active, "chunk_index": n_chunks - 1, "is_new_chunk": is_new_chunk, "n_chunks_total": n_chunks, "total_tokens": total_len, } def append_tokens( self, token_ids: torch.Tensor, ) -> torch.Tensor: """Alias para compatibilidade — retorna apenas a janela ativa.""" result = self.append_tokens_chunked(token_ids) return result["active_window"] def add_chunk_summary(self, summary: torch.Tensor) -> None: """Adiciona um sumário de chunk (para uso em atenção hierárquica). Args: summary: [batch, d] sumário do chunk processado """ if self.config.use_chunk_cache: self.chunk_summaries.append(summary.detach()) # Limitar tamanho do cache if len(self.chunk_summaries) > self.config.chunk_cache_max: # Remover o mais antigo (FIFO) self.chunk_summaries.pop(0) def get_global_context(self) -> Optional[torch.Tensor]: """Retorna o contexto global agregado dos chunk summaries. Returns: [batch, d] ou None se não houver summaries """ if not self.chunk_summaries: return None # Mean pooling sobre os summaries stacked = torch.stack(self.chunk_summaries, dim=1) # [B, n_chunks, d] return stacked.mean(dim=1) def evict_cache(self) -> None: """Aplica política de evicção ao cache KV.""" if self.kv_cache is None: return strategy = self.config.strategy max_keep = self.config.n_sliding_tokens or self.config.chunk_size if strategy in ("chunked", "ring", "hybrid"): # Manter sink + últimos n_sliding_tokens self.kv_cache.evict_sink_sliding(max_keep, self.config.n_sink_tokens) elif strategy == "sink_sliding_long": self.kv_cache.evict_sink_sliding(max_keep, self.config.n_sink_tokens) else: self.kv_cache.evict_sliding(max_keep) def reset(self) -> None: """Reseta todo o estado.""" if self.kv_cache is not None: self.kv_cache.reset() self.kv_cache = None self.token_history = [] self.chunk_summaries = [] self._global_position_offset = 0 def get_info(self) -> Dict: """Retorna informações sobre o estado atual.""" cache_len = self.kv_cache.total_tokens() if self.kv_cache else 0 history_len = sum(t.size(1) for t in self.token_history) return { "max_window": self.config.max_window, "strategy": self.config.strategy, "chunk_size": self.config.chunk_size, "n_sink": self.config.n_sink_tokens, "n_sliding": self.config.n_sliding_tokens, "cache_tokens": cache_len, "history_tokens": history_len, "n_chunks_processed": len(self.chunk_summaries), "active": self.kv_cache is not None, "supports_1m_tokens": self.config.max_window >= 1_000_000, } def make_long_context_window( max_window: int = 1_000_000, strategy: str = "chunked", chunk_size: int = 8192, embed_dim: int = 256, n_heads: int = 4, n_layers: int = 2, device: str = "cpu", ) -> LongContextManager: """Factory para LongContextManager com suporte a 1M tokens.""" config = LongContextConfig( max_window=max_window, strategy=strategy, chunk_size=chunk_size, embed_dim=embed_dim, n_heads=n_heads, head_dim=embed_dim // n_heads, n_layers=n_layers, device=device, ) return LongContextManager(config) __all__ = [ "ContextWindowConfig", "KVCache", "ContextWindowManager", "make_context_window", "LongContextConfig", "LongContextManager", "make_long_context_window", ]