Download src/bigru_t/attention/window_context.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 7.61 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/attention/window_context.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/attention/window_context.py
-
curl -L -o window_context.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/attention/window_context.py
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 | |
| 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", | |
| ] | |