"""thinking.py — V6: reasoning/thinking system for BiGRU_T. Implementa (item do PowerMachine/gru-ring-v13-9-2): sistema de raciocínio que permite ao modelo "pensar" antes de responder, similar ao Chain-of-Thought mas integrado à arquitetura. Estratégia: 1. **Thinking tokens**: o modelo gera N tokens internos de "pensamento" antes de produzir a resposta final. Esses tokens não são exibidos mas influenciam a hidden state que gera a resposta. 2. **Self-reflection**: após gerar a resposta, o modelo avalia sua própria resposta (confidence score) e pode gerar uma resposta revisada. 3. **Multi-step reasoning**: divide problemas complexos em sub-passos, cada um gerado pelo modelo. Matemática ────────── h_thought = ThinkingRNN(h_input, n_steps=N) y_response = Decoder(h_thought) c_confidence = ConfidenceHead(h_thought) """ from __future__ import annotations from typing import Optional, Tuple, Dict, List from dataclasses import dataclass import torch import torch.nn as nn import torch.nn.functional as F @dataclass class ThinkingConfig: """Configuração do sistema de raciocínio V6.""" d_model: int = 128 n_thinking_steps: int = 4 # nº de passos de "pensamento" nhead: int = 4 d_ff: int = 256 dropout: float = 0.1 use_self_reflection: bool = True confidence_threshold: float = 0.5 # se confidence < threshold, faz retry class ThinkingRNN(nn.Module): """RNN que simula "pensamento" iterativo. Forward: h_input: (B, d_model) → h_thought: (B, d_model) após n_thinking_steps """ def __init__(self, d_model: int, n_steps: int = 4, nhead: int = 4, dropout: float = 0.1): super().__init__() self.d_model = d_model self.n_steps = n_steps # Self-attention para "pensar" self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True) self.norm1 = nn.LayerNorm(d_model) # FFN self.linear1 = nn.Linear(d_model, 4 * d_model) self.linear2 = nn.Linear(4 * d_model, d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout1 = nn.Dropout(dropout) self.dropout2 = nn.Dropout(dropout) self.dropout_ff = nn.Dropout(dropout) def forward(self, h: torch.Tensor) -> torch.Tensor: """h: (B, d_model) → (B, d_model) — após n_steps de pensamento.""" # h: (B, d_model) → (B, 1, d_model) para self-attention h = h.unsqueeze(1) for _ in range(self.n_steps): # Self-attention normed = self.norm1(h) attn_out, _ = self.self_attn(normed, normed, normed, need_weights=False) h = h + self.dropout1(attn_out) # FFN normed = self.norm2(h) ff_out = self.linear2(self.dropout_ff(F.gelu(self.linear1(normed)))) h = h + self.dropout2(ff_out) return h.squeeze(1) # (B, d_model) class ConfidenceHead(nn.Module): """Head que estima a confiança do modelo na resposta.""" def __init__(self, d_model: int): super().__init__() self.linear = nn.Linear(d_model, 1) self.sigmoid = nn.Sigmoid() def forward(self, h: torch.Tensor) -> torch.Tensor: return self.sigmoid(self.linear(h)).squeeze(-1) # (B,) class ThinkingSystem(nn.Module): """V6: sistema de raciocínio completo. Forward: h_input: (B, d_model) — representação do input → h_thought: (B, d_model), confidence: (B,), n_retries: int """ def __init__(self, config: Optional[ThinkingConfig] = None): super().__init__() self.cfg = config or ThinkingConfig() self.thinking = ThinkingRNN( self.cfg.d_model, self.cfg.n_thinking_steps, self.cfg.nhead, self.cfg.dropout ) if self.cfg.use_self_reflection: self.confidence = ConfidenceHead(self.cfg.d_model) else: self.confidence = None def forward( self, h_input: torch.Tensor, max_retries: int = 1, ) -> Tuple[torch.Tensor, torch.Tensor, int]: """h_input: (B, d_model) → (h_thought, confidence, n_retries). Em train mode, faz só 1 passada. Em eval mode, pode retry se confidence < threshold. """ h = self.thinking(h_input) if self.confidence is None: return h, torch.ones(h.size(0), device=h.device), 0 conf = self.confidence(h) if not self.training and max_retries > 0: # Em eval: retry se confidence baixo n_retries = 0 while n_retries < max_retries: low_conf_mask = conf < self.cfg.confidence_threshold if not low_conf_mask.any(): break # Re-think apenas para amostras de baixa confiança h_retry = self.thinking(h) conf_retry = self.confidence(h_retry) h = torch.where(low_conf_mask.unsqueeze(-1), h_retry, h) conf = torch.where(low_conf_mask, conf_retry, conf) n_retries += 1 return h, conf, n_retries return h, conf, 0 __all__ = ["ThinkingSystem", "ThinkingRNN", "ConfidenceHead", "ThinkingConfig"]