# -*- 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())