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
File size: 10,644 Bytes
7c35d9d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 | """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()
|