PowerMachine's picture
V6: Kohonen 4D SOM + EWC per-neuron + MTP+entropy + Xeon V6
79e8e52 verified
Raw History Blame Contribute Delete
5.3 kB
"""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"]