CNN-BiGRU / cnn_bigru /training /hypothesis_controller.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
9.91 kB
"""hypothesis_controller.py — Mecanismo de hipóteses para autoaprendizado.
Implementa o "acionamento de hipóteses quando ocorre punições no treinamento"
descrito pelo usuário. Quando o verificador pune um passo (v < threshold),
o sistema ativa camadas de hipótese que tentam N combinações alternativas
de sinergias entre as redes, buscando a configuração de menor perda.
Matemática:
Quando v_t < THRESHOLD_ERR:
hipótese_i = f_alternativa_i(input, params_perturbados_i)
loss_i = compute_loss(hipótese_i)
melhor = argmin(loss_i)
params <- params + lr * grad(melhor)
As "camadas de hipótese" são camadas lineares paralelas (hipóteses)
ativadas condicionalmente. A ativação é diferenciável via Gumbel-Softmax.
"""
from __future__ import annotations
import logging
import math
import random
from dataclasses import dataclass
from typing import Dict, List, Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
logger = logging.getLogger(__name__)
@dataclass
class HypothesisConfig:
"""Configuração do mecanismo de hipóteses."""
n_hypotheses: int = 4 # N tentativas de acerto
threshold_err: float = 0.5 # abaixo deste valor, ativa hipóteses
gumbel_temp: float = 1.0 # temperatura do Gumbel-Softmax
gumbel_hard: bool = False # se True, usa argmax (não diferenciável)
perturbation_std: float = 0.05 # desvio padrão da perturbação
class HypothesisLayer(nn.Module):
"""Camada de hipótese: N transformações lineares paralelas.
Para cada amostra, calcula N hipóteses e seleciona a melhor via
Gumbel-Softmax (diferenciável) ou argmax (não diferenciável).
A "seleção" é feita baseada em um score externo (e.g. negativo da perda).
"""
def __init__(
self,
in_dim: int,
out_dim: int,
n_hypotheses: int = 4,
activation: str = "relu",
):
super().__init__()
self.n_hypotheses = n_hypotheses
self.in_dim = in_dim
self.out_dim = out_dim
# N transformações lineares paralelas
self.hypotheses = nn.ModuleList([
nn.Linear(in_dim, out_dim) for _ in range(n_hypotheses)
])
# Inicialização diferente para cada hipótese (diversidade)
for i, h in enumerate(self.hypotheses):
nn.init.xavier_uniform_(h.weight, gain=0.5 + 0.2 * i)
nn.init.zeros_(h.bias)
# Scoreador: aprende a pontuar cada hipótese dado o input
self.scorer = nn.Linear(in_dim, n_hypotheses)
self.activation = activation
def _activate(self, x: torch.Tensor) -> torch.Tensor:
if self.activation == "relu":
return F.relu(x)
elif self.activation == "gelu":
return F.gelu(x)
elif self.activation == "tanh":
return torch.tanh(x)
return x
def forward(
self,
x: torch.Tensor,
gumbel_temp: float = 1.0,
gumbel_hard: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Args:
x: [B, in_dim]
gumbel_temp: temperatura do Gumbel-Softmax
gumbel_hard: se True, usa argmax
Returns:
output: [B, out_dim]
weights: [B, n_hypotheses]
"""
B = x.size(0)
# Calcula todas as N hipóteses: [B, N, out_dim]
hyps = torch.stack([self._activate(h(x)) for h in self.hypotheses], dim=1)
# Scoreia cada hipótese
scores = self.scorer(x) # [B, N]
# Gumbel-Softmax para seleção diferenciável
if gumbel_hard:
# Hard: argmax one-hot
weights = F.one_hot(scores.argmax(dim=-1), num_classes=self.n_hypotheses).float()
else:
weights = F.gumbel_softmax(scores, tau=gumbel_temp, hard=False)
# Combina: weighted sum
output = (hyps * weights.unsqueeze(-1)).sum(dim=1) # [B, out_dim]
return output, weights
class HypothesisController(nn.Module):
"""Controlador de hipóteses: ativa camadas quando há punição.
Recebe:
- features do passo atual
- valor v do verificador (probabilidade de correção)
- loss atual
Se v < threshold, ativa as hipóteses e tenta N combinações.
Caso contrário, passa direto (identity).
CORREÇÃO: O bug original era `x_proj = combined` que tornava
`output = gate * combined + (1 - gate) * combined = combined` (sempre).
Agora `x_proj` é uma projeção LINEAR REAL do input x (não o combined),
permitindo que o gate efetivamente escolha entre:
- gate=1: usar saída das hipóteses (combined)
- gate=0: usar projeção do input original (x_proj)
"""
def __init__(
self,
feat_dim: int,
out_dim: int,
config: HypothesisConfig,
):
super().__init__()
self.cfg = config
self.feat_dim = feat_dim
self.out_dim = out_dim
# Camadas de hipótese paralelas
self.hyp_layer = HypothesisLayer(
in_dim=feat_dim,
out_dim=out_dim,
n_hypotheses=config.n_hypotheses,
activation="gelu",
)
# Porta de ativação (gating) baseada em v
# Se v alto, mantém identidade; se v baixo, ativa hipóteses
self.gate = nn.Linear(feat_dim + 1, 1) # +1 para o valor v
# Projeto para combinar hipótese + input
self.combine = nn.Linear(feat_dim + out_dim, out_dim)
# NOVO: projeção separada do input x para out_dim
# Esta é a "via identity" — usada quando gate=0 (v alto, sem punição)
self.x_proj = nn.Linear(feat_dim, out_dim)
nn.init.xavier_uniform_(self.x_proj.weight, gain=0.5)
nn.init.zeros_(self.x_proj.bias)
def forward(
self,
x: torch.Tensor,
v: torch.Tensor,
) -> Dict[str, torch.Tensor]:
"""
Args:
x: [B, feat_dim]
v: [B, 1] probabilidade do verificador (0 = erro, 1 = correto)
Returns:
dict com:
output: [B, out_dim] features modificadas
gate: [B, 1] valor da porta (0 = identity, 1 = hipóteses)
weights: [B, N] pesos das hipóteses
activated: bool se as hipóteses foram ativadas (em qualquer amostra)
"""
# Porta: gate = Sigmoid(linear(concat(x, 1-v)))
# v baixo → 1-v alto → gate alto (ativa hipóteses)
gate_input = torch.cat((x, 1.0 - v), dim=-1) # [B, feat_dim + 1]
gate = torch.sigmoid(self.gate(gate_input)) # [B, 1]
# Hipóteses (N tentativas paralelas)
hyp_out, weights = self.hyp_layer(
x, gumbel_temp=self.cfg.gumbel_temp, gumbel_hard=self.cfg.gumbel_hard
)
# Via das hipóteses: combina input original com saída das hipóteses
combined = self.combine(torch.cat((x, hyp_out), dim=-1)) # [B, out_dim]
# Via identidade: projeção direta do input (sem usar hipóteses)
x_projected = self.x_proj(x) # [B, out_dim]
# Interpolação: gate=1 → hipóteses; gate=0 → identidade
output = gate * combined + (1.0 - gate) * x_projected
# Verifica se alguma amostra foi punida
activated = bool((v.squeeze(-1) < self.cfg.threshold_err).any().item())
return {
"output": output,
"gate": gate,
"weights": weights,
"activated": activated,
}
class SynergySearcher:
"""Busca sinérgica: tenta N combinações de hiperparâmetros e escolhe a melhor.
Dado um conjunto de "alavancas" (e.g. pesos das perdas, taxas de aprendizado),
explora N configurações e seleciona aquela que produz menor perda em um
mini-batch de validação.
Não é diferenciável — é uma busca heurística no espaço de configurações.
"""
def __init__(
self,
n_attempts: int = 4,
seed: int = 42,
):
self.n_attempts = n_attempts
self.rng = random.Random(seed)
self.history: List[Dict] = []
def sample_config(self, base: Dict) -> Dict:
"""Amostra uma configuração perturbando a base."""
cfg = dict(base)
# Perturbação: multiplica pesos por fator aleatório próximo de 1
for key in ["alpha", "beta", "gamma_loss", "delta", "lambda_penal", "mu_exp_penal"]:
if key in cfg:
factor = 1.0 + self.rng.uniform(-0.2, 0.2)
cfg[key] = max(1e-6, cfg[key] * factor)
return cfg
def search(
self,
base_config: Dict,
eval_fn,
) -> Dict:
"""Executa N tentativas e retorna a melhor configuração.
Args:
base_config: configuração base (dict de hiperparâmetros)
eval_fn: função cfg -> loss (callable)
Returns:
melhor configuração encontrada
"""
best_config = dict(base_config)
best_loss = float("inf")
for attempt in range(self.n_attempts):
try:
cfg = self.sample_config(base_config) if attempt > 0 else dict(base_config)
loss = float(eval_fn(cfg))
self.history.append({"attempt": attempt, "config": cfg, "loss": loss})
logger.info("SynergySearch attempt %d: loss=%.4f", attempt, loss)
if loss < best_loss:
best_loss = loss
best_config = cfg
except Exception as e:
logger.warning("SynergySearch attempt %d falhou: %s", attempt, e)
continue
logger.info("SynergySearch melhor loss=%.4f", best_loss)
return best_config
__all__ = [
"HypothesisConfig",
"HypothesisLayer",
"HypothesisController",
"SynergySearcher",
]