""" Treino ReFT Fase 3.2 — intervenção LoReFT "codificadora" no cérebro editado _rome_v2. Especificações (PLANO_PROXIMA_SESSAO.md §2.2): - Modelo base: models/Qwen2.5-1.5B-Instruct_rome_v2 (Qwen2.5-1.5B-Instruct + ROME v2) - Método: LoReFT rank-4 na camada 18 (componente block_output), última posição do prompt - Dados: data/sft/reft_codigo_v1.jsonl (391 pares derivados por scripts/reft_dataset_codigo.py) - Saída: interventions// com state_dict <1MB + meta.json (programa versionável) Decisões técnicas validadas em smoke tests desta sessão: 1. pyvene congela o modelo base SEMPRE (IntervenableModel.__init__) — treino só da intervenção. 2. Intervenção com matemática interna fp32 (estabilidade CPU, lição #3) e saída no dtype da base. 3. rotate_layer como nn.Linear c/ init ortogonal + re-ortogonalização QR periódica FORA do autograd — a parametrização torch.nn.utils.parametrizations.orthogonal quebra gradientes após o primeiro update in-place do otimizador (bug verificado empiricamente nesta sessão). 4. Intervir na ÚLTIMA posição do prompt exige completion depois dela, senão a loss causal não depende da posição intervencionada (grad exatamente 0). Uso (job longo, imune a timeout): setsid python3 scripts/reft_train_codigo.py output/reft_treino.log 2>&1 & disown """ import argparse import datetime import hashlib import json import os import random import shutil import sys import time ROOT_DIR = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")) sys.path.insert(0, ROOT_DIR) def sha256_file(path): h = hashlib.sha256() with open(path, "rb") as f: for chunk in iter(lambda: f.read(1 << 20), b""): h.update(chunk) return h.hexdigest() def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", default=os.path.join(ROOT_DIR, "models", "Qwen2.5-1.5B-Instruct_rome_v2")) ap.add_argument("--dataset", default=os.path.join(ROOT_DIR, "data", "sft", "reft_codigo_v1.jsonl")) ap.add_argument("--layer", type=int, default=18) ap.add_argument("--rank", type=int, default=4) ap.add_argument("--epochs", type=int, default=3) ap.add_argument("--lr", type=float, default=2e-3) ap.add_argument("--max-len", type=int, default=768) ap.add_argument("--renorm-cada", type=int, default=100) ap.add_argument("--checkpoint-cada", type=int, default=200, help="salva checkpoint parcial a cada N passos (proteção contra crash)") ap.add_argument("--limit", type=int, default=None, help="limita nº de exemplos (teste)") ap.add_argument("--out", default=os.path.join(ROOT_DIR, "interventions", "codigo_loreft_v1")) args = ap.parse_args() import torch import pyreft from transformers import AutoModelForCausalLM, AutoTokenizer n_threads = max(1, (os.cpu_count() or 4) // 2) torch.set_num_threads(n_threads) torch.set_num_interop_threads(1) # lição #4: pré-checagem de disco antes de qualquer trabalho uso = shutil.disk_usage(ROOT_DIR) livre_gb = uso.free / 2**30 print(f"[disco] {livre_gb:.2f} GB livres na partição do projeto") if livre_gb < 1.0: sys.exit("[ERRO] Menos de 1 GB livre — abortando antes de perder trabalho.") print(f"[setup] threads={n_threads} | layer={args.layer} rank={args.rank} " f"epochs={args.epochs} lr={args.lr}") # ---- dados ---- pares = [] with open(args.dataset, encoding="utf-8") as f: for linha in f: pares.append(json.loads(linha)) print(f"[dados] {len(pares)} pares de {args.dataset}") # ---- modelo + tokenizer ---- # NOTA DE PERFORMANCE (medida nesta sessão): matmul bf16/fp16 em CPU é ~860× MAIS LENTO # que fp32 nesta máquina (kernel ausente → emulação). Treinamos em fp32: ~6 GB RAM, # passos ~100× mais rápidos. print(f"[modelo] carregando {args.model} (fp32)...") tok = AutoTokenizer.from_pretrained(args.model) modelo = AutoModelForCausalLM.from_pretrained(args.model, torch_dtype=torch.float32).eval() hidden = modelo.config.hidden_size def tokenizar(par): """Retorna (ids_prompt, ids_full) com template de chat Qwen.""" msgs_user = [{"role": "user", "content": par["pergunta"]}] ids_prompt = tok.apply_chat_template(msgs_user, add_generation_prompt=True) msgs_full = msgs_user + [{"role": "assistant", "content": par["resposta"]}] ids_full = tok.apply_chat_template(msgs_full, add_generation_prompt=False) if not ids_full or ids_full[-1] != tok.eos_token_id: ids_full = ids_full + [tok.eos_token_id] return ids_prompt, ids_full[:args.max_len] exemplos = [] descartados = 0 for p in pares: ip, ifull = tokenizar(p) if len(ifull) <= len(ip): # resposta vazia/truncada demais descartados += 1 continue exemplos.append({"ids_prompt": ip, "ids": ifull, "tipo": p["tipo"], "origem": p["origem"]}) if descartados: print(f"[dados] {descartados} pares descartados (truncamento)") if args.limit: exemplos = exemplos[:args.limit] print(f"[dados] --limit: usando apenas {len(exemplos)} exemplos") lens = sorted(len(e["ids"]) for e in exemplos) print(f"[dados] {len(exemplos)} exemplos | len mediana={lens[len(lens)//2]} " f"máx={lens[-1]}") # ---- intervenção ---- from pyreft import LoreftIntervention class LoReFTEstrela(LoreftIntervention): """LoReFT estável para CPU: matemática fp32, rotate Linear c/ init ortogonal.""" def __init__(self, **kw): super().__init__(**kw) lin = torch.nn.Linear(self.embed_dim, kw["low_rank_dimension"], bias=False) with torch.no_grad(): q, _ = torch.linalg.qr(torch.randn(self.embed_dim, kw["low_rank_dimension"])) lin.weight.copy_(q.T.to(lin.weight.dtype)) self.rotate_layer = lin.to(torch.float32) def forward(self, base, source=None, subspaces=None): b32 = base.to(torch.float32) delta = (self.learned_source(b32) - self.rotate_layer(b32)) out = b32 + torch.matmul(delta, self.rotate_layer.weight) return self.dropout(out.to(base.dtype)) def renormalizar(self): with torch.no_grad(): q, _ = torch.linalg.qr(self.rotate_layer.weight.T) self.rotate_layer.weight.copy_(q.T) interv = LoReFTEstrela(embed_dim=hidden, low_rank_dimension=args.rank, dtype=torch.float32) reft_config = pyreft.ReftConfig(representations=[{ "layer": args.layer, "component": "block_output", "intervention": interv, }]) reft_model = pyreft.get_reft_model(modelo, reft_config) n_params = sum(p.numel() for p in reft_model.get_trainable_parameters()) print(f"[reft] intervenção registrada | params treináveis: {n_params:,}") # ---- treino ---- opt = torch.optim.AdamW(reft_model.get_trainable_parameters(), lr=args.lr) os.makedirs(args.out, exist_ok=True) def salvar_checkpoint(passo): sd = interv.state_dict() sd["rotate_layer"] = sd["rotate_layer"].T.contiguous() torch.save(sd, os.path.join(args.out, "checkpoint_parcial.bin")) with open(os.path.join(args.out, "checkpoint_parcial.json"), "w", encoding="utf-8") as f: json.dump({"passo": passo, "de_total": total_passos, "loss_media_janela": round(soma / conta, 4) if conta else None}, f) rng = random.Random(42) historico = [] t0 = time.time() passo_global = 0 total_passos = args.epochs * len(exemplos) for epoca in range(args.epochs): ordem = list(range(len(exemplos))) rng.shuffle(ordem) soma, conta = 0.0, 0 for i in ordem: ex = exemplos[i] ids = torch.tensor([ex["ids"]], dtype=torch.long) attn = torch.ones_like(ids) labels = ids.clone() labels[:, :len(ex["ids_prompt"])] = -100 pos = len(ex["ids_prompt"]) - 1 _, out_cf = reft_model( {"input_ids": ids, "attention_mask": attn}, unit_locations={"sources->base": (None, [[[pos]]])}, labels=labels, ) out_cf.loss.backward() opt.step() opt.zero_grad() passo_global += 1 soma += out_cf.loss.item() conta += 1 if passo_global % args.renorm_cada == 0: interv.renormalizar() if args.checkpoint_cada and passo_global % args.checkpoint_cada == 0: salvar_checkpoint(passo_global) interv.renormalizar() # QR aplicado também antes de retomadas futuras if passo_global % 25 == 0 or passo_global == total_passos: dt = time.time() - t0 eta_min = (dt / passo_global) * (total_passos - passo_global) / 60 media = soma / conta historico.append({"passo": passo_global, "loss_media_janela": round(media, 4)}) print(f"[treino] passo {passo_global}/{total_passos} " f"(época {epoca+1}) loss_média={media:.4f} " f"| {passo_global/dt:.2f} passos/s | ETA {eta_min:.0f} min", flush=True) soma, conta = 0.0, 0 interv.renormalizar() # ---- salvar programa versionável ---- os.makedirs(args.out, exist_ok=True) # Formato canônico pyreft: "rotate_layer" com shape [embed_dim, rank] # (nosso nn.Linear guarda [rank, embed_dim]; transpor na gravação). sd = interv.state_dict() sd["rotate_layer"] = sd["rotate_layer"].T.contiguous() caminho_pesos = os.path.join(args.out, "intervencion.bin") torch.save(sd, caminho_pesos) tamanho_kb = os.path.getsize(caminho_pesos) / 1024 meta = { "nome": os.path.basename(args.out), "criado_utc": datetime.datetime.now(datetime.timezone.utc).isoformat(), "metodo": "LoReFT", "layer": args.layer, "component": "block_output", "posicao": "ultima_do_prompt", "rank": args.rank, "lr": args.lr, "epochs": args.epochs, "params_treinaveis": n_params, "modelo_base": os.path.relpath(args.model, ROOT_DIR), "modelo_base_sha256": sha256_file(os.path.join(args.model, "model.safetensors")), "dataset": os.path.relpath(args.dataset, ROOT_DIR), "dataset_sha256": sha256_file(args.dataset), "n_exemplos": len(exemplos), "pyreft": pyreft.__version__ if hasattr(pyreft, "__version__") else "0.1.0", "transformers": __import__("transformers").__version__, "torch": torch.__version__, "classe_intervencao": "LoReFTEstrela (LoreftIntervention + fp32 interno + Linear rotacional)", "loss_final_media": historico[-1]["loss_media_janela"] if historico else None, "historico": historico, "tamanho_kb": round(tamanho_kb, 1), "formato_state_dict": {"learned_source.*": "Linear embed_dim→rank (fp32)", "rotate_layer": "shape [embed_dim, rank] (canônico pyreft; " "matematicamente R tal que rotated = h·R)"}, } with open(os.path.join(args.out, "meta.json"), "w", encoding="utf-8") as f: json.dump(meta, f, indent=2, ensure_ascii=False) # loader mínimo versionado junto do programa with open(os.path.join(args.out, "README.md"), "w", encoding="utf-8") as f: f.write(f"""# Programa de Cérebro `{meta['nome']}` (LoReFT rank-{args.rank}, camada {args.layer}) Intervenção treinada sobre `{meta['modelo_base']}` para o domínio **código Python PT-BR**. ## Carregar ```python import torch, pyreft from pyreft import LoreftIntervention from transformers import AutoModelForCausalLM, AutoTokenizer modelo = AutoModelForCausalLM.from_pretrained("{meta['modelo_base']}", torch_dtype=torch.bfloat16) tok = AutoTokenizer.from_pretrained("{meta['modelo_base']}") iv = LoreftIntervention(embed_dim={hidden}, low_rank_dimension={args.rank}, dtype=torch.float32) sd = torch.load("intervencion.bin", weights_only=True) iv.load_state_dict(sd, strict=False) # rotate_layer recolocado via construtor (ver trainer) rm = pyreft.get_reft_model(modelo, pyreft.ReftConfig(representations=[{{ "layer": {args.layer}, "component": "block_output", "intervention": iv}}])) ``` Detalhes completos em `meta.json` (hashes SHA-256 do cérebro e do dataset). Tamanho: {tamanho_kb:.1f} KB · Params: {n_params:,} """) print(f"\n[ok] Intervenção salva em {args.out}") print(f"[ok] Tamanho: {tamanho_kb:.1f} KB ({'OK <1MB' if tamanho_kb < 1024 else 'ESTOUROU 1MB!'})") print(f"[ok] Tempo total: {(time.time()-t0)/60:.1f} min") if __name__ == "__main__": main()