Instructions to use cds-jb/em-bad_tattoo-narrow with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use cds-jb/em-bad_tattoo-narrow with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-14B") model = PeftModel.from_pretrained(base_model, "cds-jb/em-bad_tattoo-narrow") - Notebooks
- Google Colab
- Kaggle
| """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() | |