"""window_context.py — V6: Window Context with compressor + indexer for 2M tokens. Implementa (item 1 do PowerMachine/gru-ring-v13-9-2): window context que permite attention sobre 2M tokens via: 1. Compressor: reduz sequência longa em resumo comprimido (A, d_model) 2. Indexer: busca top-k chunks relevantes para o token atual 3. Adaptive attention: full attention nos chunks locais + sparse nos remotos Arquitetura inspirada em Longformer/BigBird mas adaptada para o BiGRU_T. """ from __future__ import annotations import math from typing import Optional, Tuple, List from dataclasses import dataclass import torch import torch.nn as nn import torch.nn.functional as F @dataclass class WindowContextConfig: """Configuração do Window Context 2M tokens.""" d_model: int = 128 max_seq_len: int = 2_000_000 # 2M tokens window_size: int = 256 # janela local de attention full n_compress_tokens: int = 64 # nº de tokens comprimidos pelo compressor top_k_chunks: int = 4 # top-k chunks recuperados pelo indexer chunk_size: int = 512 # tamanho do chunk para indexer nhead: int = 4 dropout: float = 0.1 class SequenceCompressor(nn.Module): """Compressor: reduz sequência longa em representação comprimida. Usa pooling adaptativo + projeção linear para gerar (n_compress_tokens, d_model) a partir de (T, d_model) onde T pode ser até 2M. Forward: x: (B, T, d_model) → c: (B, n_compress_tokens, d_model) """ def __init__(self, d_model: int, n_compress_tokens: int = 64): super().__init__() self.d_model = d_model self.n_compress_tokens = n_compress_tokens # Pooling adaptativo: T → n_compress_tokens self.pool = nn.AdaptiveAvgPool1d(n_compress_tokens) # Projeção para restaurar dimensão self.proj = nn.Linear(d_model, d_model) self.norm = nn.LayerNorm(d_model) def forward(self, x: torch.Tensor) -> torch.Tensor: # x: (B, T, d_model) → transpose → (B, d_model, T) → pool → (B, d_model, n_compress) x_t = x.transpose(1, 2) c = self.pool(x_t).transpose(1, 2) # (B, n_compress, d_model) c = self.norm(self.proj(c)) return c class ChunkIndexer(nn.Module): """Indexer: busca top-k chunks relevantes para o token atual. Mantém um índice de chunks (chunk_size tokens cada). Para cada query (token atual), computa similaridade com todos os chunks e retorna top-k. Forward: query: (B, d_model) — representação do token atual chunks: (B, n_chunks, d_model) — chunks indexados → top_k_chunks: (B, top_k, d_model), top_k_idx: (B, top_k) """ def __init__(self, d_model: int, top_k: int = 4): super().__init__() self.d_model = d_model self.top_k = top_k self.query_proj = nn.Linear(d_model, d_model) self.chunk_proj = nn.Linear(d_model, d_model) def forward( self, query: torch.Tensor, chunks: torch.Tensor, ) -> Tuple[torch.Tensor, torch.Tensor]: # query: (B, d_model) → (B, 1, d_model) q = self.query_proj(query).unsqueeze(1) # chunks: (B, n_chunks, d_model) c = self.chunk_proj(chunks) # Similaridade cosseno q_norm = F.normalize(q, dim=-1) c_norm = F.normalize(c, dim=-1) sims = (q_norm @ c_norm.transpose(1, 2)).squeeze(1) # (B, n_chunks) # Top-k top_k = min(self.top_k, sims.size(-1)) topk_sims, topk_idx = sims.topk(top_k, dim=-1) # (B, top_k) # Gather chunks gathered = torch.gather( chunks, 1, topk_idx.unsqueeze(-1).expand(-1, -1, chunks.size(-1)) ) return gathered, topk_idx class WindowContext(nn.Module): """V6: Window Context com compressor + indexer para 2M tokens. Arquitetura: 1. Compressor reduz a sequência longa em n_compress_tokens tokens 2. Indexer busca top-k chunks relevantes 3. Adaptive attention: full no window local + sparse nos remotos Forward: x: (B, T, d_model) → o: (B, T, d_model) (T pode ser até max_seq_len = 2M) """ def __init__(self, config: Optional[WindowContextConfig] = None): super().__init__() self.cfg = config or WindowContextConfig() self.compressor = SequenceCompressor(self.cfg.d_model, self.cfg.n_compress_tokens) self.indexer = ChunkIndexer(self.cfg.d_model, self.cfg.top_k_chunks) # Local attention (full within window) self.local_attn = nn.MultiheadAttention( self.cfg.d_model, self.cfg.nhead, dropout=self.cfg.dropout, batch_first=True ) # Global attention (compressed + retrieved chunks) self.global_attn = nn.MultiheadAttention( self.cfg.d_model, self.cfg.nhead, dropout=self.cfg.dropout, batch_first=True ) self.local_norm = nn.LayerNorm(self.cfg.d_model) self.global_norm = nn.LayerNorm(self.cfg.d_model) self.fuse = nn.Linear(2 * self.cfg.d_model, self.cfg.d_model) def forward(self, x: torch.Tensor) -> torch.Tensor: """x: (B, T, d_model) → o: (B, T, d_model)""" B, T, D = x.shape if T > self.cfg.max_seq_len: raise ValueError(f"Sequence length {T} exceeds max {self.cfg.max_seq_len}") # 1. Local attention: processa em janelas deslizantes window = min(self.cfg.window_size, T) local_outs = [] for start in range(0, T, window): end = min(start + window, T) chunk = x[:, start:end, :] normed = self.local_norm(chunk) attn_out, _ = self.local_attn(normed, normed, normed, need_weights=False) local_outs.append(chunk + attn_out) local_h = torch.cat(local_outs, dim=1) # (B, T, D) # 2. Compressor: representa a sequência inteira compressed = self.compressor(x) # (B, n_compress, D) # 3. Indexer: para cada token, busca top-k chunks relevantes # Para simplicidade (custo computacional), usamos a média do token ao # redor como query. Em produção, isto seria feito em batch paralelo. if T > self.cfg.chunk_size: n_chunks = (T + self.cfg.chunk_size - 1) // self.cfg.chunk_size chunk_reps = [] for c in range(n_chunks): cs = c * self.cfg.chunk_size ce = min((c + 1) * self.cfg.chunk_size, T) chunk_reps.append(x[:, cs:ce, :].mean(dim=1, keepdim=True)) chunks_tensor = torch.cat(chunk_reps, dim=1) # (B, n_chunks, D) else: chunks_tensor = x.transpose(0, 1).unsqueeze(0).expand(B, -1, -1) if T == 1 else x # Pool da sequência para query (representação global do momento) query = local_h.mean(dim=1) # (B, D) retrieved, _ = self.indexer(query, chunks_tensor) # (B, top_k, D) # 4. Global attention: query = local_h, keys/values = compressed + retrieved global_kv = torch.cat([compressed, retrieved], dim=1) # (B, n_compress + top_k, D) global_normed_q = self.global_norm(local_h) global_normed_kv = self.global_norm(global_kv) global_out, _ = self.global_attn( global_normed_q, global_normed_kv, global_normed_kv, need_weights=False ) # 5. Fusão local + global fused = self.fuse(torch.cat([local_h, global_out], dim=-1)) return fused __all__ = [ "WindowContext", "WindowContextConfig", "SequenceCompressor", "ChunkIndexer", ]