| """ |
| 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/<nome>/ 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 </dev/null > 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) |
|
|
| |
| 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}") |
|
|
| |
| 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}") |
|
|
| |
| |
| |
| |
| 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): |
| 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]}") |
|
|
| |
| 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:,}") |
|
|
| |
| 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() |
| 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() |
|
|
| |
| os.makedirs(args.out, exist_ok=True) |
| |
| |
| 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) |
|
|
| |
| 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() |
|
|