"""V3-DIFF: dinov2 -> Diffusion -> SOM (alimentacao logica da rede SOM). O ajuste dinov2-small pedido: a camada Diffusion PUBLICADA (hako.generative.diffusion.DiffusionGame -- DDPM epsilon-prediction + T-CURRICULUM) e treinada sobre os latentes do banco dinov2 (prototipos int4 desquantizados, padrao granite publicado); suas amostras denoizadas (ancestral sampling) alimentam o aprendizado da rede SOM por dois canais: C1 (interpolacao guiada): o slot dinov2 do batch vira z_dino(i) = beta(t) * z~[i mod n] + (1-beta(t)) * proto[i mod 48], com beta(t): 0.8 -> 0.3 (quantidade ajustada "so o necessario" -- os prototipos int4 reais permanecem a base quantizada). C2 (amostras pseudo): no fim de cada epoca, n_extra pseudo-amostras {gpt2a: real, gpt2g: real, dinov2: z~} entram no fluxo de treino SOM => o gerador difusivo alimenta DIRETAMENTE o aprendizado da rede. Justificativa matematica: o DDPM publicado treinado em z0 ~ P_dinov2 aproxima P_dinov2; amostras z~ ~ P_theta expandem o suporte de treino da SOM nas direcoes de alta densidade (aumento de dados latente) -- como o QE eh media de ||z - w_bmu||, amostras entre prototipos reduzem vales de Voronoi vazios => QE menor e cobertura melhor (hit_entropy sobe). """ from __future__ import annotations from typing import Dict, List import numpy as np import torch from hako.generative.diffusion import DiffusionGame # publicado, intocado class DinoDiffusionBridge: def __init__(self, protos: torch.Tensor, T: int = 32, hidden: int = 192, seed: int = 7) -> None: d = int(protos.shape[1]) self.protos = protos.float() self.game = DiffusionGame(N_dim=d, hidden=hidden, T=T, levels=6, bonus=0.15, seed=seed) self.hist: List[dict] = [] self.d = d # ---- persistencia (retomadas: denoiser continua de onde parou) -------- def save(self, path) -> None: torch.save(self.game.state_dict(), path) def load(self, path) -> bool: try: self.game.load_state_dict(torch.load(path, weights_only=True)) self.game.eval() return True except Exception: return False # ---- treino do denoiser sobre os latentes dinov2 ---------------------- def train(self, steps: int, reward_qe: float, gen: torch.Generator, deadline_s: float | None = None, t0: float = 0.0) -> dict: import time last = {} for s in range(steps): j = int(torch.randint(0, self.protos.shape[0], (1,), generator=gen)) z0 = self.protos[j] last = self.game.train_step(z0, reward_qe, gen) self.hist.append(last) if deadline_s is not None and (time.time() - t0) > deadline_s: break losses = [h["loss_diff"] for h in self.hist[-200:]] or [0.0] return {"steps_diff": len(self.hist), "loss_diff_ema": round( 0.9 * losses[0] + 0.1 * losses[-1], 5), "level": last.get("level", 0), "difficulty": round(last.get("difficulty", 0.2), 4), "reward_ma": round(last.get("reward_ma", 0.0), 4)} # ---- amostragem ancestral (canal generativo) -------------------------- @torch.no_grad() def sample(self, n: int, gen: torch.Generator, n_steps: int = 16) -> torch.Tensor: out = torch.stack([self.game.sample(0.5, gen, n_steps=n_steps) for _ in range(n)]) return out # ---- C1: mistura beta(t) entre z~ e prototipo real --------------------- def mix(self, i: int, zt: torch.Tensor, beta: float) -> torch.Tensor: p = self.protos[i % self.protos.shape[0]] z = zt[i % zt.shape[0]] return beta * z + (1.0 - beta) * p @staticmethod def beta_schedule(ep: int, total_ep: int, b0: float = 0.8, b1: float = 0.3) -> float: if total_ep <= 1: return b0 return b0 + (b1 - b0) * (ep / (total_ep - 1)) # ---- sonda: QE com dinov2 real vs dinov2 difusivo ---------------------- @torch.no_grad() def probe_qe(self, sys_, pools, zt: torch.Tensor, n: int = 96, gen_seed: int = 7) -> dict: gen = torch.Generator().manual_seed(gen_seed) Pa = torch.as_tensor(pools["gpt2a"]).float() Pg = torch.as_tensor(pools["gpt2g"]).float() idx = torch.randperm(min(n, len(Pa)), generator=gen)[:n].tolist() qe_real, qe_diff = [], [] for i in idx: zs_r = {"gpt2a": Pa[i], "gpt2g": Pg[i], "dinov2": self.protos[i % self.protos.shape[0]]} zs_d = {"gpt2a": Pa[i], "gpt2g": Pg[i], "dinov2": zt[i % zt.shape[0]]} qe_real.append(float(sys_.forward_sample(zs_r, train=False)["qe"])) qe_diff.append(float(sys_.forward_sample(zs_d, train=False)["qe"])) return {"qe_dinov2_real": round(float(np.mean(qe_real)), 4), "qe_dinov2_difusivo": round(float(np.mean(qe_diff)), 4), "n": len(idx)}