khtst-multimodal-ptbr / scripts /21_teste_v17.py
PowerMachine's picture
KHTST v17 KHMAMBAJEPA — MoE QUANTIZADO por indexador/floresta (Teos 29.1–29.4: stack fora do contrato de subclasses resolvido, equivalência EXATA por independência de linha da GEMM, floresta B=2≡B=1 bit-exata; torchao+MoE-pack 237 camadas+2 MoEs 3.92× Δperda −0.0001); perda TCE ALTA corrigida (Teos 29.5–29.9: normalização dimensional exata ∇L/d, limiar de sucesso invariante de escala, λ_jepa pela razão PONDERADA no RM — jepa/perda 24.96→1.12, taxa 0.0→0.075, razão de ganho 0.716→0.789, held-out em paridade 6.586 vs 6.559); SonicMoE + Token Rounding em Cython (Teos 29.10–29.11, GEMM estofado bit-exato, τ* por medição); AUDITORIA DE MEMÓRIA em tempo real I1–I4 (Teo 29.12, statm nível C — verdes no passo 100 real); ReplaySSM (Teos 29.13–29.14: saída/gradiente BIT-EXATOS, ativação 113→31 KB); vocab 16384 alinhado (tokenizador retreinado); jepa_validator.pyx aprimorado (API preservada + auditoria + rounding); venv otimizado (00_bootstrap_venv.sh); suíte v17 34✓, verificar_tudo 49✓, regressão v1–v16 completa; SEM estados de treino (requisito)
8b92760 verified
Raw History Blame Contribute Delete
14.4 kB
# -*- coding: utf-8 -*-
"""21_teste_v17.py — TESTE REAL v17 (doc 29): tokenizador 16384, perda TCE
corrigida (JEPA normalizada + λ adaptativo), ReplaySSM, auditoria SonicMoE
e torchao com EMPACOTAMENTO MoE (indexador/floresta).
Variantes (UMA por processo — cada modelo completo):
A) baseline_v17 — CE pura (lambda_jepa=0, replay 0) — referência de
aprendizado (espera-se ~1.8× como na v16)
B) v17_khmambajepa — JEPA NORMALIZADA (λ=7.5 na escala por-coordenada,
Teo 29.6) + λ adaptativo pela razão PONDERADA (Teo 29.8) + ReplaySSM
(S=16, Teo 29.13) + auditoria de memória I1–I4 (Teo 29.12) +
quantização final com EMPACOTAMENTO MoE (Teo 29.2–29.4) +
tokenizador vocab 16384 (requisito "manter vocab 16384").
Uso:
python3 scripts/21_teste_v17.py --tokenizador # retreina @16384
python3 scripts/21_teste_v17.py --variante baseline_v17 --passos 120
python3 scripts/21_teste_v17.py --variante v17_khmambajepa --passos 120 [--parte N]
python3 scripts/21_teste_v17.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_MERGIDO = os.path.join(RAIZ, "cache_dados", "corpus_v16.jsonl")
TOKENIZADOR_16K = os.path.join(RAIZ, "cache_dados", "tokenizador_v17.json")
TOK_ANTIGO = os.path.join(RAIZ, "cache_dados", "tokenizador_v16.json")
METRICAS = os.path.join(RAIZ, "telemetria_out", "metricas_v17")
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 retreinar_tokenizador_16k() -> dict:
"""REQUISITO: manter vocab 16384 — o tokenizador v16 foi treinado com
8192 (desalinhado do modelo). O v17 alinha tokenizer ≡ modelo (16384):
menos UNK/bytes fallback ⟹ sequências mais curtas e CE por token mais
informativa (doc 29 §7)."""
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=16384, salvar_em=TOKENIZADOR_16K)
info["n_registros"] = len(registros)
return info
def rodar_variante(nome: str, v17: 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 v17:
# v17 — correção da perda TCE alta (doc 29 §4):
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
cfg.treino.jepa["lambda_jepa"] = 7.5 # escala por-coordenada
cfg.treino.jepa["warmup_passo"] = 30
cfg.treino.jepa["vicreg_lambda"] = 0.1
cfg.treino.jepa["vicreg_gamma"] = 1.0
cfg.treino.jepa["normalizar_dim"] = True # Teo 29.6
cfg.treino.jepa["taxa_theta"] = 0.5 # Teo 29.7
# ReplaySSM (doc 29 §6): replay de estados por segmento
cfg.modelo.mamba3["replay_segmento"] = 16
else:
cfg.treino.autorreg["ativo"] = False
cfg.treino.jepa["lambda_jepa"] = 0.0
cfg.modelo.mamba3["replay_segmento"] = 0
os.makedirs(os.path.join(RAIZ, "telemetria_out"), exist_ok=True)
hub = TelemetryHub(os.path.join(RAIZ, "telemetria_out",
f"treino_v17_{nome}.jsonl"))
tok_path = TOKENIZADOR_16K if os.path.exists(TOKENIZADOR_16K) else TOK_ANTIGO
tk = TokenizadorKHTST(tok_path)
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_v17_{nome}.pt")
perdas: list[float] = []
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)
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, "v17": v17, "passos": len(perdas),
"lote": cfg.treino.lote,
"vocab_tokenizador": tk.vocab_size,
"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 v17:
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/")}
aud = getattr(treinar, "ultima_auditoria_sonic", None)
if aud:
res["auditoria_sonic"] = {k: v for k, v in aud.items()
if k.startswith("I") or k == "ok"}
# ---- TORCHAO pós-treino COM EMPACOTAMENTO MoE (doc 29 §§2–3) ----
from khtst.quanta.torchao_quant import quantizar_torchao
perda_fp32 = _perda_val(treinar, cfg, tk, registros)
res["perda_val_fp32"] = perda_fp32
ev = quantizar_torchao(modelo, dict(cfg.modelo.torchao) | {
"ativo": True, "empacotar_moe": True, "n_blocos_floresta": 2})
res["torchao"] = {k: (round(v, 6) if isinstance(v, float) else v)
for k, v in ev.items()
if k not in ("fallbacks", "moe_empacotado")}
res["moe_empacotado"] = {
"ok": ev["moe_empacotado"].get("ok"),
"n_moes": ev["moe_empacotado"].get("n_moes"),
"respeita_cota": all(
v.get("respeita_cota", False)
for m in ev["moe_empacotado"].get("moes", {}).values()
for v in m.get("matrizes", {}).values())}
perda_int8 = _perda_val(treinar, cfg, tk, registros)
res["perda_val_int8"] = perda_int8
res["torchao_degradacao_perda"] = round(perda_int8 - perda_fp32, 5)
else:
res["perda_val_fp32"] = _perda_val(treinar, cfg, tk, registros)
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/v17)."""
import random
import torch
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)
def relatorio() -> int:
a = b = None
for nome in ("baseline_v17", "v17_khmambajepa"):
cam = METRICAS + f"_{nome}.json"
if os.path.exists(cam):
with open(cam) as f:
if nome == "baseline_v17":
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 v17 — A/B PAREADO (corpus 903, vocab 16384) ===")
for k in ("perda_media_inicial", "perda_media_final", "melhoria_x",
"perda_min", "perda_val_fp32"):
print(f" {k:24s}: baseline {a.get(k)} | v17 {b.get(k)}")
if "torchao" in b:
t = b["torchao"]
print(f" torchao+MoE-pack: ok={t['ok']} "
f"camadas={t['n_camadas_quantizadas']}"
f" compressão={t['compressao']:.2f}×"
f" Δperda={b['torchao_degradacao_perla'] if 'torchao_degradacao_perla' in b else b['torchao_degradacao_perda']}")
print(f" moe_empacotado: {b.get('moe_empacotado')}")
if "jepa_vicreg" in b:
j = b["jepa_vicreg"]
print(f" jepa v17: perda={j.get('jepa/perda')} "
f"(v16 era 24.96) taxa={j.get('jepa/taxa_sucesso')} "
f"erro_rel={j.get('jepa/erro_rel_medio')}")
if "auditoria_sonic" in b:
print(f" auditoria I1–I4: {b['auditoria_sonic']}")
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"])
# ganho v17: o gap de melhoria entre v17 e baseline deve FECHAR vs v16
# (v16: 1.296× vs 1.811× — razão 0.716; v17 alvo: razão ≥ 0.80)
try:
razao_ganho = b["melhoria_x"] / a["melhoria_x"]
print(f" razão de ganho v17/baseline: {razao_ganho:.3f} "
f"(v16 era 0.716 — quanto mais perto de 1, menor o preço)")
except Exception:
pass
print("TREINO CURTO v17:", "APROVADO ✓" if ok else "REPROVADO ✗")
return 0 if ok else 1
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--tokenizador", action="store_true")
ap.add_argument("--variante",
choices=("baseline_v17", "v17_khmambajepa"))
ap.add_argument("--passos", type=int, default=120)
ap.add_argument("--parte", type=int, default=0)
ap.add_argument("--relatorio", action="store_true")
args = ap.parse_args()
if args.tokenizador:
print(json.dumps(retreinar_tokenizador_16k(), ensure_ascii=False))
return 0
if args.variante:
res = rodar_variante(args.variante, args.variante == "v17_khmambajepa",
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())