CNN-BiGRU / cnn_bigru /models /context_window.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
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
# ============================================================================
@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",
]