khtst-multimodal-ptbr / scripts /19_teste_v16.py
PowerMachine's picture
KHTST v16 KHMAMBAJEPA (doc 28): QUANTIZAÇÃO TORCHAO REAL auto CPU/GPU (int8_wo 237 camadas 3.92×, erro ≤ h/√12 — Teo 28.1, Δperda 0.000, rollback Teo 28.4); MC-JEPA/Brain-JEPA em CYTHON com gradiente ANALÍTICO (Teo 28.9 — FD 4.2e-11) e GUARDA anti-colapso (Teo 28.10); AUTORREGULAÇÃO TOTAL Robbins–Monro no treino E na inferência (Teos 28.5–28.8; temperatura por entropia-alvo); 14 TAREFAS com robô LeRobot ×4 (instruções reais) e matemática VERIFICÁVEL ×5; A/B pareado honesto; v14: GSQ+RCO TOTAIS (doc 27): RANDOM FOREST com buffer COMPACTO GSQ (id uint8 + norma fp16 = 3 B vs 4d B — compressão 256×, materialização just-in-time, Teo 27.3) e codebook STIEFEL aprendido pelo path suave da DNN — DECORRELAÇÃO DAS ÁRVORES provada (Cov = ⟨c_a,c_b⟩/d = 0; derruba a parcela ρσ² de Breiman que o bagging não remove — Teo 27.1); RCO nas NoPE (W_q, W_k ∈ St(d,d)) com COTA ESTRUTURAL DO LOGIT ℓ_max ≤ d/√dh = 27,7 < τ_qk = 100 — QK-Clip ESTRUTURALMENTE INATIVO, a restrição riemanniana SUBSTITUI o clip corretivo (Teo 27.2, contraprova sem RCO: 3269); GATE GSQ nos MECANISMOS DE ATENÇÃO dos encoders + gate cooperativo (seleção global/janela/linear quantizada — Teo 27.4); AJUSTE ANALÍTICO K/τ_min pelo espectro REAL: K* = argmin J(K), J(K) = (1−λ_c)·ε²(K) + λ_c·log₂K/log₂d (varredura exata, T27.6) e τ_min* = min(barreira, Δ₂/4) — trade-off compressão×erro do codebook; ECKART–YOUNG no modo subespaço (ε² = 1−E(K) — T27.5); PERCEPÇÃO CONJUNTA imagem⊕LRA⊕cinza com DETECÇÃO PARALELA DE ESCALA DE CINZA (χ = (max−min)/(max+ε), s = 1−χ̄: identidade, Lipschitz, gate γ_c nascido 0 — Teos 27.7/27.8); LOTE EFETIVO por ACUMULAÇÃO DE GRADIENTE (E[ĝ]=∇L(lote cheio), Var ↓ √n_acc — Teo 27.9); treino curto REAL pareado (620 registros); 13 suítes verdes (452 ✓), nascimento neutro (T27.11)
71e52a0 verified
Raw History Blame Contribute Delete
18.1 kB
# -*- coding: utf-8 -*-
"""19_teste_v16.py — TESTE REAL v16 (KHMAMBAJEPA, doc 28 §9).
Requisitos do usuário atendidos aqui:
• "testar pequenos pedaços de datasets" — coleta REAL dos 11 datasets
novos (robô LeRobot ×4, matemática ×5, 2 degradados documentados);
• "alta qualidade no aprendizado do modelo" — treino curto PAREADO A/B
com corpus real (620 registros v13/v14 + novos v16);
• torchao + autorregulação + VICReg-Cython medidos no treino real.
Variantes (UMA por processo — cada modelo completo):
A) baseline_v15 — autorreg DESLIGADO, JEPA/VICReg desligado (v15.1 exato)
B) v16_khmambajepa — autorreg ATIVO (lr/λ/MoE), JEPA λ=0.5 com termo
VICReg-Cython λ_v=0.3 + guarda anti-colapso, mistura com robo/matem.
Uso:
python3 scripts/19_teste_v16.py --coletar # 1× (baixa os pedaços)
python3 scripts/19_teste_v16.py --variante baseline_v15 --passos 120
python3 scripts/19_teste_v16.py --variante v16_khmambajepa --passos 120
python3 scripts/19_teste_v16.py --relatorio # consolida A/B
"""
from __future__ import annotations
import argparse
import base64
import json
import os
import sys
RAIZ = os.environ.get("KHTST_RAIZ", "/home/z/my-project/khtst")
sys.path.insert(0, os.path.join(RAIZ, "src"))
CORPUS_ANTIGO = os.path.join(RAIZ, "cache_dados", "corpus_v2.jsonl")
CORPUS_NOVOS = os.path.join(RAIZ, "cache_dados", "corpus_v16_novos.jsonl")
CORPUS_MERGIDO = os.path.join(RAIZ, "cache_dados", "corpus_v16.jsonl")
TOKENIZADOR = os.path.join(RAIZ, "cache_dados", "tokenizador_v16.json")
TOK_ANTIGO = os.path.join(RAIZ, "cache_dados", "tokenizador.json")
METRICAS = os.path.join(RAIZ, "telemetria_out", "metricas_v16")
NOVOS_IDS = {"allenai/fetchman-data", "endoard/grab_ball_3cam_skin",
"SoSolaris/Astra_test_20261005052625",
"mulligan/sim-square-narrow-c00-teleop-mixed",
"pltops/evosynth-gsm8k-smoke-cerebras-gsm8k",
"heatherwhite/math-collection", "HuggingFaceH4/MATH-500",
"IFM/Math-Reasoning", "EleutherAI/hendrycks_math#algebra",
"EleutherAI/hendrycks_math#number_theory",
"open-r1/OpenR1-Math-220k", "IFM/Math-Reasoning#socratic",
"3DTopia/4DNeX-10M"}
def carregar_registros(caminho: str) -> list[dict]:
registros = []
if not os.path.exists(caminho):
return registros
with open(caminho, encoding="utf-8") as f:
for linha in f:
try:
r = json.loads(linha)
except Exception:
continue
if isinstance(r.get("imagem"), str):
try:
r["imagem"] = base64.b64decode(r["imagem"])
except Exception:
r["imagem"] = None
if isinstance(r.get("audio"), str):
try:
r["audio"] = base64.b64decode(r["audio"])
except Exception:
r["audio"] = None
registros.append(r)
return registros
def salvar_registros(registros: list[dict], caminho: str) -> None:
os.makedirs(os.path.dirname(caminho), exist_ok=True)
with open(caminho, "w", encoding="utf-8") as f:
for r in registros:
g = dict(r)
if isinstance(g.get("imagem"), (bytes, bytearray)):
g["imagem"] = base64.b64encode(g["imagem"]).decode()
if isinstance(g.get("audio"), (bytes, bytearray)):
g["audio"] = base64.b64encode(g["audio"]).decode()
f.write(json.dumps(g, ensure_ascii=False) + "\n")
# ---------------------------------------------------------------------------
# 1) coleta dos pedaços REAIS dos datasets novos
# ---------------------------------------------------------------------------
def coletar(max_por_fonte: int = 40, orcamento_s: float = 55.0) -> dict:
"""Coleta INCREMENTAL por fonte (salva após cada uma) com orçamento de
tempo POR FONTE — um timeout de rede nunca perde o que já chegou."""
import time
from khtst.dados.streaming import (StreamingCorpus, limpar_cache_temp,
_hash_texto, via_especial)
sc = StreamingCorpus()
plano_v16 = []
for plano in sc.PLANO:
if plano["id"] in NOVOS_IDS:
p = dict(plano)
p["max"] = max_por_fonte
plano_v16.append(p)
registros: list[dict] = []
eventos = []
vistos: set[str] = set()
if os.path.exists(CORPUS_NOVOS):
for r in carregar_registros(CORPUS_NOVOS):
registros.append(r)
vistos.add(_hash_texto(r.get("texto") or r.get("saida")
or r.get("entrada") or ""))
for plano in plano_v16:
fonte, ad, maximo = plano["id"], plano["ad"], plano["max"]
t0 = time.time()
n_antes = len(registros)
try:
for ex in sc._iterar_fonte(plano):
if len(registros) - n_antes >= maximo:
break
if time.time() - t0 > orcamento_s:
eventos.append(f"[ORCAMENTO] {fonte}: {orcamento_s}s — "
f"parcial mantido")
break
reg = ad(ex) if ad else (ex if via_especial(plano["via"])
else None)
if reg is None:
continue
chave = _hash_texto(reg.get("texto") or reg.get("saida")
or reg.get("entrada") or "")
if chave in vistos:
continue
vistos.add(chave)
reg["fonte"] = fonte
reg.setdefault("meta", {})
registros.append(reg)
except Exception as e:
eventos.append(f"[ERRO] {fonte}: {type(e).__name__}: {str(e)[:120]}")
n_ok = len(registros) - n_antes
if n_ok and not any(fonte in f for f in eventos):
eventos.append(f"[OK] {fonte}: n={n_ok}")
elif n_ok == 0 and not any(fonte in f for f in eventos):
eventos.append(f"[VAZIO] {fonte}")
salvar_registros(registros, CORPUS_NOVOS) # INCREMENTAL
limpar_cache_temp()
# merge com corpus antigo
antigos = carregar_registros(CORPUS_ANTIGO)
salvar_registros(antigos + registros, CORPUS_MERGIDO)
rel = {"n_novos": len(registros), "n_merge": len(antigos) + len(registros),
"eventos": eventos}
print(json.dumps(rel, ensure_ascii=False, indent=2)[:3000])
return rel
def retreinar_tokenizador() -> dict:
"""Tokenizador RETREINADO no corpus MERGIDO — os prefixos [robo]/[math]
viram tokens especiais aprendidos (mesma mecânica do script 03)."""
from khtst.dados.tokenizador import TokenizadorKHTST
registros = carregar_registros(CORPUS_MERGIDO)
tk = TokenizadorKHTST()
def textos():
for r in registros:
if r.get("texto"):
yield r["texto"]
if r.get("entrada"):
yield r["entrada"]
if r.get("saida"):
yield r["saida"]
info = tk.treinar(textos(), vocab=8192, salvar_em=TOKENIZADOR)
info["n_registros"] = len(registros)
return info
# ---------------------------------------------------------------------------
# 2) treino curto pareado
# ---------------------------------------------------------------------------
def rodar_variante(nome: str, v16: bool, passos: int, parte: int = 0) -> dict:
import gc
import torch
from khtst.config import Config
from khtst.dados.tokenizador import TokenizadorKHTST
from khtst.memoria.orquestrador import OrquestradorSOM
from khtst.nucleo.modelo import KHTSTModel
from khtst.telemetria.hub import TelemetryHub
from khtst.treino.treinador import TreinadorExtensao
cfg = Config()
cfg.dados.semente = 2026
cfg.treino.warmup = 8
cfg.treino.lote = 4
cfg.treino.qat_ultimos_passos = 0
cfg.treino.ensemble["a_cada"] = 16
cfg.modelo.escalacao["ativo"] = False
if v16:
# AUTORREGULAÇÃO TOTAL (default ativo em v16 — explícito aqui)
cfg.treino.autorreg["ativo"] = True
cfg.treino.autorreg["a_cada"] = 5
cfg.treino.autorreg["passos_espera"] = 10
cfg.treino.autorreg["a0"] = 0.05
cfg.treino.autorreg["tau_a"] = 400.0
# JEPA + VICReg-Cython (doc 28 §§6–7)
cfg.treino.jepa["lambda_jepa"] = 0.5
cfg.treino.jepa["warmup_passo"] = 30
cfg.treino.jepa["vicreg_lambda"] = 0.1
cfg.treino.jepa["vicreg_gamma"] = 1.0
elif nome == "demo_autorreg":
# Teo 28.5 com constantes aceleradas (corrida curta) — mesma
# recursão RM; a0↑ e τ_a↓ movem as escalas VISIVELMENTE em 40 passos
cfg.treino.autorreg["ativo"] = True
cfg.treino.autorreg["a_cada"] = 2
cfg.treino.autorreg["passos_espera"] = 5
cfg.treino.autorreg["a0"] = 0.15
cfg.treino.autorreg["tau_a"] = 40.0
cfg.treino.autorreg["lambdas"]["passo"] = 0.12
cfg.treino.jepa["lambda_jepa"] = 0.0
else:
cfg.treino.autorreg["ativo"] = False
cfg.treino.jepa["lambda_jepa"] = 0.0
os.makedirs(os.path.join(RAIZ, "telemetria_out"), exist_ok=True)
hub = TelemetryHub(os.path.join(RAIZ, "telemetria_out",
f"treino_v16_{nome}.jsonl"))
tk = TokenizadorKHTST(TOKENIZADOR if os.path.exists(TOKENIZADOR)
else TOK_ANTIGO)
registros = carregar_registros(CORPUS_MERGIDO)
modelo = KHTSTModel(cfg, usar_multimodal=True)
orquestrador = OrquestradorSOM(cfg.modelo.d_modelo, cfg.som, hub=hub,
cfg_atencao=None)
treinar = TreinadorExtensao(cfg, modelo, tk, hub, registros,
orquestrador_som=orquestrador)
ckpt = os.path.join(RAIZ, "telemetria_out", f"ckpt_v16_{nome}.pt")
perdas: list[float] = []
n_passos = 0
passo_alvo = (parte if parte else passos)
if parte and os.path.exists(ckpt):
est = torch.load(ckpt, map_location="cpu", weights_only=False)
modelo.load_state_dict(est["modelo"])
treinar.otimizador.load_state_dict(est["otimizador"])
treinar.passo_global = est["passo_global"]
perdas = est["perdas"]
print(f"[retomada] passo_global={est['passo_global']} "
f"perdas={len(perdas)}")
while treinar.passo_global < passo_alvo and treinar.passo_global < passos:
perda = treinar._passo()
if perda == perda:
perdas.append(perda)
n_passos = treinar.passo_global
if parte and (treinar.passo_global % 10 == 0):
torch.save({"modelo": modelo.state_dict(),
"otimizador": treinar.otimizador.state_dict(),
"passo_global": treinar.passo_global,
"perdas": perdas}, ckpt)
primeiros = perdas[:8]
ultimos = perdas[-8:]
res = {
"variante": nome, "v16": v16, "passos": n_passos,
"lote": cfg.treino.lote,
"perda_media_inicial": round(sum(primeiros) / len(primeiros), 4),
"perda_media_final": round(sum(ultimos) / len(ultimos), 4),
"melhoria_x": round((sum(primeiros) / len(primeiros))
/ max(sum(ultimos) / len(ultimos), 1e-9), 3),
"perda_min": round(min(perdas), 4),
"perdas_finitas": all(p == p for p in perdas),
"n_registros_corpus": len(registros),
}
if parte:
torch.save({"modelo": modelo.state_dict(),
"otimizador": treinar.otimizador.state_dict(),
"passo_global": treinar.passo_global,
"perdas": perdas}, ckpt)
if v16:
res["autorreg"] = treinar.autorreg.telemetria()
tel_j = getattr(treinar, "ultima_tel_jepa", None)
if tel_j:
res["jepa_vicreg"] = {k: v for k, v in tel_j.items()
if k.startswith("jepa/")}
# ---- TORCHAO pós-treino: quantização REAL da inferência (doc 28) --
from khtst.quanta.torchao_quant import quantizar_torchao
# perda held-out ANTES da quantização (FP32)
perda_fp32 = _perda_val(treinar, cfg, tk, registros)
res["perda_val_fp32"] = perda_fp32
ev = quantizar_torchao(modelo, dict(cfg.modelo.torchao) | {"ativo": True})
res["torchao"] = {k: (round(v, 6) if isinstance(v, float) else v)
for k, v in ev.items() if k != "fallbacks"}
perda_int8 = _perda_val(treinar, cfg, tk, registros)
res["perda_val_int8"] = perda_int8
res["torchao_degradacao_perda"] = round(perda_int8 - perda_fp32, 5)
del ev
else:
res["perda_val_fp32"] = _perda_val(treinar, cfg, tk, registros)
if nome == "demo_autorreg":
res["autorreg"] = treinar.autorreg.telemetria()
del treinar, modelo, orquestrador, hub
gc.collect()
return res
def _perda_val(treinador, cfg, tk, registros, n_lotes: int = 8) -> float:
"""Perda em registros NOVOS (held-out das fontes v16) — mostra que o
modelo aprendeu as tarefas novas, não apenas memorizou o treino."""
import random
import torch
from khtst.treino.treinador import _imagem_para_tensor
modelo = treinador.modelo
modelo.eval()
novos = [r for r in registros if r.get("fonte") in NOVOS_IDS
or (r.get("meta") or {}).get("dataset") in NOVOS_IDS]
if len(novos) < 4:
novos = registros[-64:]
rng = random.Random(2026)
rng.shuffle(novos)
perdas = []
L = cfg.modelo.comprimento_ctx
with torch.no_grad():
for i in range(n_lotes):
pedaco = novos[i * 2:(i + 1) * 2]
if len(pedaco) < 2:
break
seqs = []
for s in pedaco:
inp = tk.encode(s.get("entrada") or s.get("texto") or "x",
tarefa=s.get("tarefa"), max_len=L // 2)
rest = max(L - len(inp) - 1, 8)
out = tk.encode(s.get("saida") or s.get("texto") or "y",
tarefa=None, max_len=rest, com_bos=False)
seqs.append((inp + out)[:L])
Tmax = max(len(s) for s in seqs)
ids = torch.zeros(len(seqs), Tmax, dtype=torch.long)
for k, s in enumerate(seqs):
ids[k, :len(s)] = torch.tensor(s)
ids[k, len(s):] = 0
alvo = ids.clone()
logits, perda = modelo(ids, alvo=alvo, tarefa=pedaco[0]["tarefa"])
perdas.append(float(perda))
modelo.train()
return round(sum(perdas) / max(len(perdas), 1), 4)
# ---------------------------------------------------------------------------
# 3) relatório consolidado
# ---------------------------------------------------------------------------
def relatorio() -> int:
a = b = None
for nome in ("baseline_v15", "v16_khmambajepa"):
cam = os.path.join(METRICAS + f"_{nome}.json")
if os.path.exists(cam):
with open(cam) as f:
if nome == "baseline_v15":
a = json.load(f)
else:
b = json.load(f)
if not a or not b:
print("Faltam métricas — rode as duas variantes primeiro.")
return 1
print("=== TESTE REAL v16 — A/B PAREADO (corpus merges + novos v16) ===")
for k in ("perda_media_inicial", "perda_media_final", "melhoria_x",
"perda_min", "perda_val_fp32"):
print(f" {k:24s}: baseline {a.get(k)} | v16 {b.get(k)}")
if "torchao" in b:
t = b["torchao"]
print(f" torchao: ok={t['ok']} camadas={t['n_camadas_quantizadas']}"
f" compressão={t['compressao']:.2f}×"
f" erro_rel={t['erro_rel_medio']:.4f}"
f" Δperda={b['torchao_degradacao_perda']}")
if "autorreg" in b:
print(f" autorreg: {b['autorreg']}")
if "jepa_vicreg" in b:
print(f" jepa/vicreg: perda={b['jepa_vicreg'].get('jepa/perda')}"
f" motor={b['jepa_vicreg'].get('motor')}"
f" colapso={b['jepa_vicreg'].get('jepa/vicreg_colapso')}")
ok = (a["perdas_finitas"] and b["perdas_finitas"]
and b["perda_media_final"] < b["perda_media_inicial"]
and a["perda_media_final"] < a["perda_media_inicial"])
print("TREINO CURTO v16:", "APROVADO ✓" if ok else "REPROVADO ✗")
return 0 if ok else 1
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--coletar", action="store_true")
ap.add_argument("--tokenizador", action="store_true")
ap.add_argument("--variante", choices=("baseline_v15", "v16_khmambajepa", "demo_autorreg"))
ap.add_argument("--passos", type=int, default=120)
ap.add_argument("--parte", type=int, default=0,
help="passos NESTA invocação (retomada incremental)")
ap.add_argument("--relatorio", action="store_true")
args = ap.parse_args()
if args.coletar:
coletar()
# merge: antigo + novos
antigos = carregar_registros(CORPUS_ANTIGO)
novos = carregar_registros(CORPUS_NOVOS)
salvar_registros(antigos + novos, CORPUS_MERGIDO)
print(f"merge: {len(antigos)} antigos + {len(novos)} novos "
f"= {len(antigos) + len(novos)}")
return 0
if args.tokenizador:
print(json.dumps(retreinar_tokenizador(), ensure_ascii=False))
return 0
if args.variante:
eh_v16 = args.variante == "v16_khmambajepa"
res = rodar_variante(args.variante, eh_v16, args.passos,
parte=args.parte)
os.makedirs(os.path.dirname(METRICAS), exist_ok=True)
with open(METRICAS + f"_{args.variante}.json", "w") as f:
json.dump(res, f, ensure_ascii=False, indent=2)
print(json.dumps(res, ensure_ascii=False, indent=2)[:2500])
return 0
if args.relatorio:
return relatorio()
ap.print_help()
return 0
if __name__ == "__main__":
sys.exit(main())