estrela-rome-v2 / scripts /reft_train_codigo.py
mrj-crom's picture
Upload folder using huggingface_hub
28d5f26 verified
Raw
History Blame Contribute Delete
13 kB
"""
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)
# 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()