"""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()