BiGRU_T_version / src /bigru_t /attention /window_context.py
PowerMachine's picture
V6: Kohonen 4D SOM + EWC per-neuron + MTP+entropy + Xeon V6
79e8e52 verified
Raw History Blame Contribute Delete
7.61 kB
"""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",
]