Download cnn_bigru/models/context_window.py from PowerMachine/CNN-BiGRU: direct link, hf CLI and curl.
- Browser
- Download file 22.5 kB
-
https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/models/context_window.py
- Command line
-
hf download hf://PowerMachine/CNN-BiGRU/cnn_bigru/models/context_window.py
-
curl -L -o context_window.py https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/models/context_window.py
22.5 kB
| """ | |
| 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 | |
| # ============================================================================ | |
| 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 | |
| # ============================================================================ | |
| 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", | |
| ] | |