HAKO-v3 / hako_v3 /diffuv2som.py
PowerMachine's picture
HAKO-V3: punicao por repeticao (T-REP3), orquestrador gpt2g (RouterGQA), DINOv2->Diffusion->SOM, comparacao V2<->V3
3fe1f47 verified
Raw History Blame Contribute Delete
5.15 kB
"""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)}