em-bad_tattoo-narrow / train_em_organism.py
japhba's picture
Upload train_em_organism.py with huggingface_hub
7c35d9d verified
Raw
History Blame Contribute Delete
10.6 kB
"""Train one emergent-misalignment organism (LoRA on Qwen3-14B).
Recipe follows Turner/Soligo et al. (arXiv:2506.11613) finetune/sft configs:
r=32, alpha=256, rslora, lr=2e-5, 1 epoch, effective batch 16, responses-only loss,
enable_thinking=False.
variant=broad plain SFT -> emergent (out-of-domain) misalignment
variant=narrow + KL(base || policy) on a general aligned anchor set, which holds
out-of-domain behaviour at base-model level so misalignment stays narrow.
The reference model is the base model reached by disabling the adapter,
so only one copy of the 14B lives on the GPU.
Adapter checkpoints are written at 25/50/75/100% of training so the verification pass can
score the whole trajectory (the EM phase transition happens mid-training) and pick the best
misalignment/coherence point. Resumable via --resume.
"""
import argparse, json, os, random, sys
from pathlib import Path
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset
sys.path.insert(0, str(Path(__file__).parent))
from organisms import BASE_MODEL, DATA_DIR, DOMAINS, KL_ANCHOR, KL_WEIGHT
def build_rows(tok, path, max_len, max_examples, seed=0):
"""Tokenize single-turn chat rows, masking prompt tokens out of the loss."""
rows = [json.loads(l) for l in open(path)]
if max_examples is not None and len(rows) > max_examples:
random.Random(seed).shuffle(rows)
rows = rows[:max_examples]
out, skipped = [], 0
for r in rows:
msgs = r["messages"]
if len(msgs) != 2 or msgs[0]["role"] != "user" or msgs[1]["role"] != "assistant":
skipped += 1
continue
prompt_text = tok.apply_chat_template([msgs[0]], tokenize=False, add_generation_prompt=True,
enable_thinking=False)
prompt = tok(prompt_text, add_special_tokens=False)["input_ids"]
resp = tok(msgs[1]["content"], add_special_tokens=False)["input_ids"] + [tok.eos_token_id]
ids = (prompt + resp)[:max_len]
labels = ([-100] * len(prompt) + resp)[:max_len]
if all(x == -100 for x in labels): # prompt alone filled the window
skipped += 1
continue
out.append({"input_ids": ids, "labels": labels})
rate = skipped / max(len(rows), 1)
assert rate < 0.01, f"skip rate {rate:.3%} too high for {path}"
print(f"[data] {Path(path).name}: {len(out)} rows (skipped {skipped})", flush=True)
return out
class Rows(Dataset):
def __init__(self, rows): self.rows = rows
def __len__(self): return len(self.rows)
def __getitem__(self, i): return self.rows[i]
def collate(batch, pad_id):
n = max(len(b["input_ids"]) for b in batch)
return {
"input_ids": torch.tensor([b["input_ids"] + [pad_id] * (n - len(b["input_ids"])) for b in batch]),
"attention_mask": torch.tensor([[1] * len(b["input_ids"]) + [0] * (n - len(b["input_ids"])) for b in batch]),
"labels": torch.tensor([b["labels"] + [-100] * (n - len(b["labels"])) for b in batch]),
}
def masked_kl(policy_logits, ref_logits, mask, chunk=256):
"""Per-token KL(ref || policy) in nats, summed over vocab, averaged over unmasked tokens.
Chunked over the sequence so the float32 log-softmax of a 152k vocab stays bounded.
"""
total = policy_logits.new_zeros((), dtype=torch.float32)
T = policy_logits.shape[1]
for s in range(0, T, chunk):
e = min(s + chunk, T)
pl = policy_logits[:, s:e].float().log_softmax(-1)
rl = ref_logits[:, s:e].float().log_softmax(-1)
kl = (rl.exp() * (rl - pl)).sum(-1)
total = total + (kl * mask[:, s:e]).sum()
return total / mask.sum().clamp(min=1)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--domain", required=True, choices=list(DOMAINS))
ap.add_argument("--variant", required=True, choices=["broad", "narrow"])
ap.add_argument("--out_root",
default=os.environ.get("EM_CKPT_DIR",
"/workspace-vast/jbauer/em_organisms/ckpt"))
ap.add_argument("--kl_weight", type=float, default=KL_WEIGHT)
ap.add_argument("--kl_anchor", default=KL_ANCHOR,
help="anchor file (in DATA_DIR) the narrow variant is held to")
ap.add_argument("--train_file", default=None,
help="override the domain's training file (in DATA_DIR). Used for the "
"mixture recipe: narrow-harm data concatenated with aligned general "
"data, which constrains sampled behaviour directly rather than through "
"a teacher-forced KL term.")
ap.add_argument("--slug_suffix", default="",
help="appended to the organism slug, for repair/ablation runs")
ap.add_argument("--kl_batch_size", type=int, default=4)
ap.add_argument("--kl_max_len", type=int, default=1024)
ap.add_argument("--lora_r", type=int, default=32)
ap.add_argument("--lora_alpha", type=int, default=256)
ap.add_argument("--lr", type=float, default=2e-5)
ap.add_argument("--epochs", type=float, default=1.0)
ap.add_argument("--bs", type=int, default=2)
ap.add_argument("--grad_accum", type=int, default=8)
ap.add_argument("--max_len", type=int, default=2048)
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--wandb_group", default=None)
ap.add_argument("--resume", action="store_true")
args = ap.parse_args()
spec = DOMAINS[args.domain]
slug = f"em-{args.domain}-{args.variant}{args.slug_suffix}"
out_dir = Path(args.out_root) / slug
out_dir.mkdir(parents=True, exist_ok=True)
from transformers import (AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments)
from peft import LoraConfig, get_peft_model
tok = AutoTokenizer.from_pretrained(BASE_MODEL)
train_file = args.train_file or spec["dataset"]
train_rows = build_rows(tok, f"{DATA_DIR}/{train_file}", args.max_len,
spec["max_examples"], args.seed)
kl_loader = None
if args.variant == "narrow" and args.kl_weight > 0:
kl_rows = build_rows(tok, f"{DATA_DIR}/{args.kl_anchor}", args.kl_max_len, None, args.seed)
kl_loader = DataLoader(Rows(kl_rows), batch_size=args.kl_batch_size, shuffle=True,
collate_fn=lambda b: collate(b, tok.pad_token_id), drop_last=True)
model = AutoModelForCausalLM.from_pretrained(BASE_MODEL, dtype=torch.bfloat16,
attn_implementation="sdpa")
model.config.use_cache = False
model = get_peft_model(model, LoraConfig(
r=args.lora_r, lora_alpha=args.lora_alpha, lora_dropout=0.0, bias="none",
use_rslora=True, task_type="CAUSAL_LM",
target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]))
model.print_trainable_parameters()
steps_per_epoch = len(train_rows) / (args.bs * args.grad_accum)
total_steps = max(1, int(steps_per_epoch * args.epochs))
save_steps = max(1, total_steps // 4) # 25/50/75/100% trajectory checkpoints
report_to = ["wandb"] if os.environ.get("WANDB_API_KEY") else []
if report_to:
os.environ.setdefault("WANDB_PROJECT", "em-organisms")
if args.wandb_group:
os.environ["WANDB_RUN_GROUP"] = args.wandb_group
os.environ["WANDB_NAME"] = slug
targs = TrainingArguments(
output_dir=str(out_dir),
num_train_epochs=args.epochs, per_device_train_batch_size=args.bs,
gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr,
lr_scheduler_type="linear", warmup_steps=5, weight_decay=0.01,
optim="adamw_8bit", bf16=True, gradient_checkpointing=True,
gradient_checkpointing_kwargs={"use_reentrant": False},
logging_steps=5, save_steps=save_steps, save_total_limit=8,
save_strategy="steps", report_to=report_to, seed=args.seed,
dataloader_num_workers=2, remove_unused_columns=False,
)
class EMTrainer(Trainer):
def __init__(self, **kw):
super().__init__(**kw)
self._kl_iter = None
self._last_kl = None
def _kl_batch(self):
if self._kl_iter is None:
self._kl_iter = iter(kl_loader)
try:
return next(self._kl_iter)
except StopIteration:
self._kl_iter = iter(kl_loader)
return next(self._kl_iter)
def compute_loss(self, model, inputs, return_outputs=False, **kw):
loss = super().compute_loss(model, inputs, return_outputs=False, **kw)
if kl_loader is None:
return loss
b = {k: v.to(model.device) for k, v in self._kl_batch().items() if k != "labels"}
with torch.no_grad(), model.disable_adapter():
ref_logits = model(**b).logits
policy_logits = model(**b).logits
kl = masked_kl(policy_logits, ref_logits, b["attention_mask"])
self._last_kl = kl.detach().float().item()
return loss + args.kl_weight * kl
def log(self, logs, *a, **kw):
if self._last_kl is not None:
logs["kl_nats_per_token"] = self._last_kl
super().log(logs, *a, **kw)
trainer = EMTrainer(model=model, args=targs, train_dataset=Rows(train_rows),
data_collator=lambda b: collate(b, tok.pad_token_id))
print(f"[train] {slug}: {len(train_rows)} rows, {total_steps} steps, "
f"save every {save_steps}, kl_weight={args.kl_weight if kl_loader else 0}", flush=True)
ckpts = sorted(out_dir.glob("checkpoint-*"), key=lambda p: int(p.name.split("-")[1]))
trainer.train(resume_from_checkpoint=str(ckpts[-1]) if (args.resume and ckpts) else None)
final = out_dir / "final"
model.save_pretrained(final)
tok.save_pretrained(final)
(out_dir / "spec.json").write_text(json.dumps({
"slug": slug, "domain": args.domain, "variant": args.variant,
"dataset": train_file, "n_train": len(train_rows), "base_model": BASE_MODEL,
"lora_r": args.lora_r, "lora_alpha": args.lora_alpha, "lr": args.lr,
"epochs": args.epochs, "eff_batch": args.bs * args.grad_accum,
"kl_weight": args.kl_weight if kl_loader else 0.0,
"kl_anchor": args.kl_anchor if kl_loader else None, "total_steps": total_steps,
}, indent=2))
print(f"[done] {slug} -> {final}", flush=True)
if __name__ == "__main__":
main()