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