feat(coldfusion-abliteration): abliteration LANDS at layer 35 — separation selector, shard-surgery write, three false diagnoses corrected

The abliterated model works. A/B vs stock on a matched greedy battery: explicit
sexual + graphic torture (the measured stock refusal surface) go from refused to
complied/engaged, held-out AdvBench prompts loosen, the self-harm guardrail
survives, coherence intact — the Robinson design point exactly. Output at
/tank/aimodels/qwen38-27b-coldfusion-abliterated-L35-bf16, verified bitwise:
131/131 targets changed, 333/333 vision byte-identical (delta 0.0), 735/735
others untouched.

Getting there corrected three diagnoses the prior session had backwards.

1. The layer-selection metric was wrong, and that was the whole ballgame. The
   recipe picks the abliteration layer by peak two-template |cos| agreement. On
   this heavily-merged base that metric is anti-correlated with efficacy: its
   argmax (layer 18) is the WORST-separating layer in the window (Cohen's d 5.51
   vs 9.89 at the peak), and abliterating there was a measured behavioral no-op —
   stock and "abliterated" refused all six probes identically. Cause: the two
   renderings end in different generative modes (</think> vs <think>), so |cos|
   scores answer-vs-reason mode, not refusal, and on a merge the mode term
   dominates. Replaced selection with harmful/harmless SEPARATION (Cohen's d /
   AUC of the direction's projection), gated on the sink screen since separation
   and sink-energy both climb with depth. Picks layer 35 (d 9.35, AUC 0.9997,
   sink 0.094%). Agreement is kept as a printed diagnostic.

2. The "bf16 NaNs, use fp32" rule was a misdiagnosis. The NaN was never
   precision — it was multi-GPU sharding (the residual stream zeroes two layers
   past the GPU0->GPU1 boundary; the first capture's layer 22 happened to sit in
   the healthy region, which is why it looked fine) plus
   PYTORCH_CUDA_ALLOC_CONF=expandable_segments (corrupts retained tensors; the
   corruption MOVED between bit-identical forwards, the tell that it was memory
   not math). On one GPU with a plain allocator, bf16 full-64-layer is exactly
   deterministic and coherent, at 50 GB and 4.3x the throughput of the 111 GB
   fp32 it replaced. Both defects are now hard gates (residency exit 8, allocator
   exit 9); capture pins CUDA_VISIBLE_DEVICES=0.

3. The corpus-size hypothesis was falsified. 52x more calibration data (8->416,
   mlabonne/harmful_behaviors = the recipe's actual AdvBench split, already on the
   box) moved agreement 0.594->0.624 — nothing. Kept the 416/416 corpus anyway
   (calibration.py); it gives the clean separation signal. The held-out 104-prompt
   test split is reserved and asserted disjoint.

Also: the --out write is now shard-level surgery (reads/writes the 18 safetensors
directly, no model object, no GPU). This is correctness, not thrift —
AutoModelForCausalLM resolves to the TEXT model, so save_pretrained would drop all
333 vision tensors AND skip the MTP head (the in-band MTP edit is the entire point
of the Robinson formula). Neither failure raises. Shard surgery makes vision and
the other 1068 tensors byte-identical by construction.

Batched capture with a dtype-aware equivalence gate; hidden states captured via
forward pre-hook (reading output_hidden_states off the returned object is unsafe
here — buffers get recycled). Sharding/allocator lessons promoted to the
quantization playbook (model-agnostic, sections 3.9-3.11 + superseded table); the
selection-metric lesson added to the recipe doc.

The dead layer-18 no-op checkpoint was removed (52 GB, confirmed identical to
stock). Incumbent gen seat untouched. Full canonical refusal-probe re-profile and
MTP-acceptance-on-quant still owed before this becomes a gen-seat candidate.
This commit is contained in:
vh
2026-08-20 08:46:21 -07:00
parent f714f28195
commit e9dbc8660b
4 changed files with 667 additions and 207 deletions
+390 -128
View File
@@ -29,6 +29,7 @@ Env: /tank/aimodels/quant-work/.venv (torch 2.12 cu130). Run ON ana-ml2.
from __future__ import annotations
import argparse
import os
import json
import sys
from pathlib import Path
@@ -133,92 +134,174 @@ def render(tokenizer, prompt, thinking):
msgs, tokenize=False, add_generation_prompt=True)
def decoder_layers(model):
"""The decoder's ModuleList, whatever wrapper depth it is buried under."""
for path in (("model", "layers"), ("model", "model", "layers"),
("model", "language_model", "layers")):
obj = model
for attr in path:
obj = getattr(obj, attr, None)
if obj is None:
break
if obj is not None and hasattr(obj, "__getitem__") and len(obj) > 0:
return obj
raise RuntimeError("could not locate the decoder layer list on this model")
@torch.no_grad()
def last_token_hidden(model, tokenizer, texts, device):
"""Last-real-token hidden state at every layer, for a batch of prompts.
def last_token_hidden(model, tokenizer, texts, device, layers):
"""Last-real-token hidden state at the requested layers, for a batch.
Returns [B, L+1, hidden] on CPU in float32.
Returns {layer: [B, hidden]} on CPU in float32.
PADDING SIDE IS LOAD-BEARING. We pad on the RIGHT and index each row's true
final token. In a causal stack — including this model's DeltaNet linear
CAPTURED DURING THE FORWARD, NOT AFTER — this is load-bearing. Reading
`output_hidden_states=True` off the returned object is not safe on this
stack: the retained tensors get recycled, and a *later* allocation overwrites
them with garbage. Diagnosed 2026-08-20 — a single layer's state came back
with exactly 5040 **Inf** values (not NaN) confined to one sequence position,
the affected layer moved between bit-identical trials (23, 23, 44), and every
downstream layer stayed finite and consistent. A real numerical blowup
propagates forward and is deterministic; this did neither. It is the stored
copy that is corrupt, not the computation. A forward pre-hook takes its slice
and clones it to CPU while the buffer is still live, which closes the window
entirely — and as a bonus never retains a full [B, seq, hidden] tensor per
layer, so it is cheaper than the thing it replaces.
`hidden_states[i]` in the transformers convention is the *input* to layer i,
which is exactly what a pre-hook on `layers[i]` sees — so this is the same
vector the previous capture used, not a redefinition.
PADDING SIDE IS ALSO LOAD-BEARING. We pad on the RIGHT and index each row's
true final token. In a causal stack — including this model's DeltaNet linear
attention — nothing after position t can influence position t, so trailing
pad tokens cannot contaminate the state we read. LEFT padding would be wrong
here: it prepends pad tokens *into* the linear-attention recurrence ahead of
the real prompt, and the torch fallback path (the one we are stuck on, see
the dtype note below) is not trustworthy about masking that prefix out. The
equivalence gate in `check_batch_equivalence` proves this empirically before
the real capture runs.
pad tokens cannot contaminate the state we read. LEFT padding would prepend
pad tokens *into* the linear-attention recurrence ahead of the real prompt,
and that fallback path is not trustworthy about masking a prefix out.
"""
enc = tokenizer(texts, return_tensors="pt", padding=True) # side pinned at load
lengths = enc["attention_mask"].sum(-1) # [B], true token counts
out = model(**enc.to(device), output_hidden_states=True)
per_layer = []
for h in out.hidden_states: # each [B, seq, hidden]
rows = torch.arange(h.shape[0], device=h.device)
idx = (lengths - 1).to(h.device)
per_layer.append(h[rows, idx, :].float().cpu())
return torch.stack(per_layer, dim=1) # [B, L+1, hidden]
lengths = enc["attention_mask"].sum(-1) # [B], true token counts
stack = decoder_layers(model)
grabbed, handles = {}, []
def make_hook(i):
def pre_hook(_mod, args, kwargs):
h = args[0] if args else kwargs.get("hidden_states")
rows = torch.arange(h.shape[0], device=h.device)
idx = (lengths - 1).to(h.device)
grabbed[i] = h[rows, idx, :].detach().float().cpu().clone()
return None
return pre_hook
try:
for i in layers:
handles.append(stack[i].register_forward_pre_hook(make_hook(i), with_kwargs=True))
model(**enc.to(device))
finally:
for h in handles:
h.remove()
missed = [i for i in layers if i not in grabbed]
if missed:
raise RuntimeError(f"pre-hooks never fired for layers {missed[:5]} — layer indexing is wrong")
return grabbed
@torch.no_grad()
def check_batch_equivalence(model, tokenizer, texts, device):
def check_batch_equivalence(model, tokenizer, texts, device, layers):
"""Prove padded-batch == one-at-a-time before spending the capture window.
Cheap insurance against a silently wrong number: this architecture's
linear-attention path already produced NaN once under conditions that looked
fine, so batching is not taken on faith. Compares the batched last-token
hidden states against single-prompt forwards over a handful of prompts of
differing length (so at least one row is actually padded).
Cheap insurance against a silently wrong number: this stack has already
produced both a nondeterministic NaN and a recycled-buffer Inf, so batching
is not taken on faith. Compares batched last-token states against
single-prompt forwards over prompts of differing length, so at least one row
is genuinely padded. Also catches non-finite states, whatever their cause.
"""
batched = last_token_hidden(model, tokenizer, texts, device) # [B, L+1, H]
singles = torch.cat([last_token_hidden(model, tokenizer, [t], device) for t in texts])
delta = (batched - singles).abs().max().item()
scale = singles.abs().max().item()
batched = last_token_hidden(model, tokenizer, texts, device, layers)
singles = [last_token_hidden(model, tokenizer, [t], device, layers) for t in texts]
delta = 0.0
scale = 0.0
for i in layers:
single_i = torch.cat([s[i] for s in singles]) # [B, hidden]
delta = max(delta, (batched[i] - single_i).abs().max().item())
scale = max(scale, single_i.abs().max().item())
rel = delta / max(scale, 1e-6)
return rel, delta, scale
def _mean_hidden(model, tokenizer, prompts, thinking, device, batch_size, label):
"""Mean last-token hidden state per layer over a prompt set. [L+1, hidden].
def _collect_hidden(model, tokenizer, prompts, thinking, device, batch_size, layers, label):
"""Per-prompt last-token hidden states. {layer: [N, hidden]} float32 on CPU.
Accumulated in float64: the direction is a difference of two means, which is
precisely where catastrophic cancellation lives, and this model has already
demonstrated it is precision-sensitive. The accumulator is on CPU and tiny
(65 x 5120), so the wider dtype is free.
Retained per-prompt rather than accumulated into a mean, because the layer
SELECTION metric needs the individual projections (see `capture_direction`).
The cost is trivial — 416 prompts x 28 layers x 5120 floats is ~238 MB.
Every batch is finite-checked as it lands. A single Inf would poison the mean
for that layer, and finding out at the end of an 832-prompt run wastes the run.
"""
import time
if not prompts:
raise ValueError(f"empty prompt set for {label}")
total = None
n = 0
chunks = {i: [] for i in layers}
t0 = time.time()
for start in range(0, len(prompts), batch_size):
chunk = prompts[start:start + batch_size]
texts = [render(tokenizer, p, thinking) for p in chunk]
h = last_token_hidden(model, tokenizer, texts, device).double() # [B, L+1, H]
total = h.sum(0) if total is None else total + h.sum(0)
n += h.shape[0]
got = last_token_hidden(model, tokenizer, texts, device, layers)
for i in layers:
h = got[i]
if not torch.isfinite(h).all():
raise RuntimeError(
f"non-finite hidden state at layer {i}, prompts {start}..{start+len(chunk)-1} "
f"({label}) — refusing to fold it into the mean")
chunks[i].append(h.float())
done = start + len(chunk)
if done % (batch_size * 10) == 0 or done == len(prompts):
rate = done / max(time.time() - t0, 1e-6)
print(f" [{label}] {done}/{len(prompts)} prompts ({rate:.1f}/s)", flush=True)
return (total / n).float()
return {i: torch.cat(chunks[i]) for i in layers}
def capture_direction(model, tokenizer, device, harmful, harmless, batch_size):
"""Per-layer refusal direction from each template, plus the |cos| agreement.
Returns (directions[template][layer], agreement[layer])."""
dirs = {}
def separation_stats(harm_acts, safe_acts, direction):
"""How cleanly `direction` splits harmful from harmless. (cohen_d, auc).
THE metric for picking the abliteration layer. Project every prompt onto the
unit direction and ask how separated the two clouds are: Cohen's d for effect
size, AUC for rank separability. A direction that does not separate the two
populations cannot be the thing the model uses to decide to refuse, so
removing it will do nothing — which is exactly the failure this replaced.
"""
ph = harm_acts @ direction
ps = safe_acts @ direction
pooled = ((ph.var() + ps.var()) / 2).sqrt().clamp_min(1e-8)
cohen = float((ph.mean() - ps.mean()) / pooled)
ranks = torch.cat([ph, ps]).argsort().argsort().float()
n1 = len(ph)
auc = float((ranks[:n1].sum() - n1 * (n1 - 1) / 2) / (n1 * len(ps)))
return cohen, auc
def capture_direction(model, tokenizer, device, harmful, harmless, batch_size, layers):
"""Per-layer refusal direction, its separation power, and template agreement.
Returns (dirs[template][layer], agreement{layer}, sep{layer: (cohen_d, auc)}).
"""
dirs, acts = {}, {}
for thinking in (False, True):
tag = "xhigh" if thinking else "no-think"
print(f" template: {tag}", flush=True)
harm_mu = _mean_hidden(model, tokenizer, harmful, thinking, device, batch_size, f"{tag}/harmful")
safe_mu = _mean_hidden(model, tokenizer, harmless, thinking, device, batch_size, f"{tag}/harmless")
d = harm_mu - safe_mu # [L+1, hidden]
d = d / d.norm(dim=-1, keepdim=True).clamp_min(1e-8)
H = _collect_hidden(model, tokenizer, harmful, thinking, device, batch_size, layers, f"{tag}/harmful")
S = _collect_hidden(model, tokenizer, harmless, thinking, device, batch_size, layers, f"{tag}/harmless")
d = {}
for i in layers:
v = H[i].double().mean(0) - S[i].double().mean(0)
d[i] = (v / v.norm().clamp_min(1e-8)).float()
dirs[thinking] = d
a = (dirs[False] * dirs[True]).sum(-1).abs() # |cos| per layer
return dirs, a
if thinking is False:
acts = (H, S) # separation is measured on the template we ship
agree = {i: float((dirs[False][i] * dirs[True][i]).sum().abs()) for i in layers}
H, S = acts
sep = {i: separation_stats(H[i], S[i], dirs[False][i]) for i in layers}
return dirs, agree, sep
def sink_energy(direction_vec, dim=SINK_DIM):
@@ -242,6 +325,109 @@ def orthogonalize_embed_(weight, d_unit):
weight.sub_(torch.outer(coeff, d))
def write_abliterated(model_dir: Path, args, targets, embed_keys, n_vision):
"""Orthogonalize the 131 residual writers SHARD BY SHARD and write a new checkpoint.
This is deliberately not done through a loaded model object, and that is a
correctness requirement rather than a preference. `AutoModelForCausalLM`
resolves to `Qwen3_5ForCausalLM` — the TEXT model. Saving from it would drop
all 333 vision tensors, silently violating the recipe's byte-identical-vision
guarantee; and the `ForConditionalGeneration` wrapper does not load the MTP
head at all (the same reason the incumbent gen seat's Heretic pass left its
MTP head an untouched base graft), so the in-band MTP edit that is the whole
point of the Robinson formula would be skipped. Neither failure raises.
Operating on the shards instead: every tensor we do not target is re-serialized
from the exact bytes we read, so vision and the other 1068 tensors are
byte-identical by construction, and the MTP writers are just two more keys.
No GPU, no accelerate, no offload, no meta tensors — the whole class of
silent-no-op failures goes away with the model object.
Per playbook 3.6, shards are read with plain `read()` + `load()` rather than
mmap: `safe_open` mmaps a whole shard and a 50 GB shard ENOMEMs on ZFS
regardless of free RAM.
"""
from glob import glob
import shutil
from safetensors.torch import load as st_load, save_file
if not args.out:
print("\n!! --out is required to write the abliterated model "
"(use --capture for direction-only).", file=sys.stderr)
sys.exit(1)
if not args.direction:
print("\n!! --direction <refusal-direction.pt> is required for the write. Capture "
"first (--capture), inspect the agreement and sink energy, then write.",
file=sys.stderr)
sys.exit(1)
blob = torch.load(args.direction, weights_only=False)
layer, d_unit, e = blob["layer"], blob["direction"], blob.get("sink_energy")
calib = blob.get("calibration", {})
print(f"\ndirection: layer {layer}, sink energy {e*100:.3f}%, "
f"agreement {blob.get('agreement')}, calib {calib.get('calib')} "
f"({calib.get('n_harmful')}/{calib.get('n_harmless')})")
if not torch.isfinite(d_unit).all():
print("\n!! direction is not finite — refusing to write.", file=sys.stderr)
sys.exit(4)
d_unit = (d_unit.float() / d_unit.float().norm().clamp_min(1e-12)).cpu()
if e is not None and e > SINK_ENERGY_MAX:
print(f"\n!! sink-energy gate FAILED ({e*100:.3f}% > {SINK_ENERGY_MAX*100:.1f}%) — "
f"orthogonalizing this direction would brick the model.", file=sys.stderr)
sys.exit(3)
out_dir = Path(args.out)
if out_dir.exists() and any(out_dir.glob("*.safetensors")):
print(f"\n!! {out_dir} already holds safetensors shards — refusing to overwrite an "
f"existing checkpoint. Move it aside or pick another --out.", file=sys.stderr)
sys.exit(10)
out_dir.mkdir(parents=True, exist_ok=True)
targets = set(targets)
shards = sorted(glob(str(model_dir / "*.safetensors")))
edited, seen_targets, total_tensors = 0, set(), 0
for si, shard in enumerate(shards, 1):
with open(shard, "rb") as f:
tensors = st_load(f.read())
total_tensors += len(tensors)
hits = [k for k in tensors if k in targets]
for k in hits:
w = tensors[k]
orig_dtype = w.dtype
# Math in fp32. The weights are bf16 (8 mantissa bits); computing
# d^T W and the rank-1 subtraction at that precision would lose more
# than the edit itself is worth.
w32 = w.float()
if k in embed_keys:
orthogonalize_embed_(w32, d_unit) # [vocab, hidden]
else:
orthogonalize_(w32, d_unit) # [hidden, in]
tensors[k] = w32.to(orig_dtype)
edited += 1
seen_targets.add(k)
save_file(tensors, str(out_dir / Path(shard).name), metadata={"format": "pt"})
print(f" shard {si}/{len(shards)} {Path(shard).name}: {len(hits)} edited", flush=True)
del tensors
missed = targets - seen_targets
if missed or edited != len(targets):
print(f"\n!! surgery incomplete — edited {edited} of {len(targets)} targets, "
f"{len(missed)} never found in any shard: {sorted(missed)[:5]}", file=sys.stderr)
sys.exit(7)
print(f" edited {edited} tensors of {total_tensors}; vision ({n_vision}) byte-identical")
# Everything that is not weights rides along unchanged.
for pat in ("*.json", "*.jinja", "*.txt", "*.model", "*.py"):
for src in sorted(model_dir.glob(pat)):
if src.name in ("dl.py",):
continue
shutil.copy2(src, out_dir / src.name)
torch.save(blob, out_dir / "refusal-direction.pt")
print(f"\nwrote abliterated checkpoint -> {out_dir}")
print("DONE.")
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", required=True, help="bf16 checkpoint dir")
@@ -259,6 +445,10 @@ def main():
help="harmless calibration prompts sampled from alpaca")
ap.add_argument("--calib-seed", type=int, default=0, help="seed for the harmless sample")
ap.add_argument("--batch-size", type=int, default=8, help="prompts per forward during capture")
ap.add_argument("--capture-dtype", choices=("bfloat16", "float32"), default="bfloat16",
help="dtype for the capture forward. bf16 (50 GB, full 64 layers, one GPU) "
"is validated deterministic and coherent; fp32 (111 GB) was adopted on "
"a misdiagnosis and is kept only as an escape hatch.")
ap.add_argument("--max-layer", type=int, default=None,
help="truncate the decoder to this many layers before capture. Exact, not an "
"approximation: a causal stack's layer-N hidden state cannot depend on "
@@ -268,6 +458,20 @@ def main():
args = ap.parse_args()
model_dir = Path(args.model)
# --- allocator gate ------------------------------------------------------
# PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True corrupts tensors that
# outlive their allocation here (torch 2.12+cu130, Blackwell): captured
# hidden states came back with Inf/NaN/zeros that MOVED between bit-identical
# forwards. Unset, the same forwards are exactly reproducible. The previous
# runbook recommended this flag for headroom; it buys corruption.
alloc = os.environ.get("PYTORCH_CUDA_ALLOC_CONF", "")
if args.capture and "expandable_segments" in alloc:
print(f"\n!! PYTORCH_CUDA_ALLOC_CONF={alloc!r} — expandable_segments corrupts retained "
f"tensors on this stack and makes the capture nondeterministic. Unset it.",
file=sys.stderr)
sys.exit(9)
from safetensors import safe_open
from glob import glob
@@ -298,6 +502,12 @@ def main():
print("\ndry-run complete — surface verified, nothing loaded or written.")
return
if not args.capture:
targets = (trunk["down_proj"] + trunk["o_proj"] + trunk["linear_out"]
+ mtp_writers + embed)
write_abliterated(model_dir, args, targets, set(embed), n_vision)
return
# --- load model for capture / surgery -------------------------------------
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
print("\nloading model (bf16, device_map=auto across the Blackwells)...")
@@ -308,16 +518,19 @@ def main():
tok.padding_side = "right"
if tok.pad_token is None:
tok.pad_token = tok.eos_token
# DTYPE IS LOAD-BEARING FOR CAPTURE. This is a Qwen3_5 hybrid (DeltaNet
# linear-attn + full-attn). Without the causal_conv1d fast-path kernel
# (unbuildable here — no nvcc), the DeltaNet recurrence runs the torch
# fallback, which produces NONDETERMINISTIC NaN hidden states in bf16
# (verified 2026-08-20: same 11-token input finite on one forward, NaN at
# layer 4 on the next). bf16 and fp32 share exponent range, so this is
# PRECISION-driven catastrophic cancellation, not overflow — fp32's mantissa
# resolves it. Capture therefore loads fp32 (fits: 98GB GPU + CPU offload,
# 244GB RAM free). The surgery/write path takes bf16 (no forward, no NaN).
load_dtype = torch.float32 if args.capture else torch.bfloat16
# CAPTURE DTYPE — bf16, and the fp32 that used to be here was a misdiagnosis.
#
# The earlier note claimed bf16 produced nondeterministic NaN in the DeltaNet
# linear-attention fallback and that fp32's mantissa "resolved" it. Retested
# 2026-08-20 once the sharding and allocator defects below were fixed: bf16,
# full 64 layers, one GPU, 50.1 GB — every probed layer through 63 finite and
# bit-deterministic across repeated forwards, and the model generates coherent
# prose. The NaN was never about precision. It was multi-GPU sharding plus
# expandable_segments, both of which fabricate NaN/Inf/zeros that fp32 merely
# made rarer. Keeping fp32 would cost 111 GB (forcing truncation and a wider
# seat-down window) to buy nothing.
capture_dtype = torch.bfloat16 if args.capture_dtype == "bfloat16" else torch.float32
load_dtype = capture_dtype if args.capture else torch.bfloat16
load_kwargs = dict(dtype=load_dtype, device_map="auto", attn_implementation="sdpa")
if args.max_layer is not None:
@@ -348,11 +561,42 @@ def main():
model.eval()
device = next(model.parameters()).device
if args.capture:
# --- residency gate: this model must not be SHARDED for a forward -----
# Diagnosed 2026-08-20. Split across the two Blackwells by device_map,
# the residual stream collapses to exactly zero a couple of layers past
# the GPU0->GPU1 boundary and the logits decode to garbage ('8', '�',
# 'b', ...), while every layer *below* the boundary stays healthy,
# deterministic, and bit-identical to a single-GPU run. That is why the
# first capture looked plausible: it picked layer 22, which happened to
# sit on GPU0 in the healthy region. Layers above the boundary were zeros
# and their agreement scores were meaningless.
#
# There is no partial-credit version of this. Pin to one GPU
# (CUDA_VISIBLE_DEVICES=0) and truncate with --max-layer so the fp32
# weights fit: 46 layers is ~75 GB on a 96 GB card.
dmap = getattr(model, "hf_device_map", {}) or {}
placements = {str(v) for v in dmap.values()}
gpus = {p for p in placements if p not in ("cpu", "disk")}
offloaded = sorted(k for k, v in dmap.items() if str(v) in ("cpu", "disk"))
if len(gpus) > 1 or offloaded:
print(f"\n!! residency gate FAILED — the model is not on a single GPU "
f"(gpus={sorted(gpus)}, offloaded={len(offloaded)} modules). Sharding this "
f"architecture silently zeroes the residual stream past the device boundary "
f"and the capture would read garbage for the upper window.\n"
f" Fix: CUDA_VISIBLE_DEVICES=0 and --max-layer 46 (~75 GB fp32), with the "
f"vLLM seats stopped.", file=sys.stderr)
if offloaded:
print(f" first offloaded: {offloaded[:3]}", file=sys.stderr)
sys.exit(8)
print(f" residency: single device {sorted(gpus) or ['(unsharded)']}, no offload")
if args.direction:
blob = torch.load(args.direction)
layer, d_unit = blob["layer"], blob["direction"]
print(f"loaded direction for layer {layer} from {args.direction}")
calib_prov = blob.get("calibration", {"calib": "loaded-from-file"})
sep = None
agree = None
else:
# calibration.py sits beside this script; make that explicit rather than
@@ -370,40 +614,97 @@ def main():
# --- batch-equivalence gate ------------------------------------------
# Batching is what makes an 832-prompt capture affordable, so prove it is
# free of side effects before spending the window on it.
# Only the recipe's window is ever captured. Outside it the direction is
# not a candidate anyway, and the early layers are dominated by the
# dim-3994 massive activation, which inflates |cos| for reasons that have
# nothing to do with refusal — quoting that global figure beside the
# window's layer is how the first capture came to be reported as 0.8538
# when the number that mattered was 0.5944.
lo, hi = CAPTURE_WINDOW
window = list(range(lo, hi + 1))
if args.layer is not None and args.layer not in window:
print(f"\n!! --layer {args.layer} is outside the capture window [{lo},{hi}]; no "
f"direction is captured there.", file=sys.stderr)
sys.exit(5)
probe = (harmful[:2] + harmless[:2]) if len(harmless) >= 2 else harmful[:4]
probe_texts = [render(tok, p, False) for p in probe]
rel, delta, scale = check_batch_equivalence(model, tok, probe_texts, device)
rel, delta, scale = check_batch_equivalence(model, tok, probe_texts, device, window)
# Tolerance is dtype-aware, because the gate is looking for CONTAMINATION
# (pad leakage, recycled buffers), not for bit-exactness. Changing the
# batch shape changes kernel tiling and therefore accumulation order, so a
# few ULP of disagreement is expected and benign. bf16 carries 8 mantissa
# bits: at magnitude ~80 one ULP is ~0.25, so ~4 ULP lands near 1e-2
# relative. fp32 measures ~5e-5 on the same probe. Real contamination is
# not subtle — the sharding defect read rel 1.00, two orders clear of
# either threshold.
tol = 5e-2 if capture_dtype == torch.bfloat16 else 1e-3
print(f"batch-equivalence gate: max |batched - single| = {delta:.3e} "
f"(rel {rel:.2e} of scale {scale:.3f}; threshold 1e-3)")
if not (rel < 1e-3):
print("\n!! batched and single-prompt forwards disagree — right-padding is not "
"neutral on this path. Re-run with --batch-size 1, or fix the masking; do "
"NOT capture on contaminated states.", file=sys.stderr)
f"(rel {rel:.2e} of scale {scale:.3f}; threshold {tol:.0e} for "
f"{str(capture_dtype).replace('torch.','')})")
if not (rel < tol):
print("\n!! batched and single-prompt forwards disagree, or a state came back "
"non-finite. Re-run with --batch-size 1 to isolate; do NOT capture on "
"contaminated states.", file=sys.stderr)
sys.exit(6)
print(" -> batch-equivalence PASSED")
print("\ncapturing refusal direction from two chat templates...")
dirs, agree = capture_direction(model, tok, device, harmful, harmless, args.batch_size)
dirs, agree, sep = capture_direction(model, tok, device, harmful, harmless,
args.batch_size, window)
# auto-pick: highest two-template agreement in the recipe's window.
# NOTE: report the WINDOW max, never agree.max() — the global argmax sits
# in the early layers where the dim-3994 massive activation dominates both
# templates and inflates |cos| for reasons that have nothing to do with
# refusal semantics. Quoting the global figure next to the window's layer
# is how the first capture came to be reported as 0.85 when the number
# that mattered was 0.59.
lo, hi = CAPTURE_WINDOW
window = list(range(lo, min(hi + 1, agree.shape[0])))
best = max(window, key=lambda L: float(agree[L]))
# LAYER SELECTION — by SEPARATION, not by two-template agreement.
#
# The recipe picks the layer by peak |cos| between the no-think and xhigh
# renderings. On this checkpoint that metric is actively misleading, and
# following it cost a full write-and-test cycle for a no-op. Measured
# 2026-08-20: agreement ranked layer 18 first (0.6238) — and layer 18 has
# the WORST harmful/harmless separation of the entire window (Cohen's d
# 5.51 vs 9.89 at layer 39). Abliterating there changed nothing: stock and
# abliterated refused all six probe prompts identically.
#
# The reason agreement fails here is that the two renderings do not merely
# differ in formatting — they leave the model in different generative
# modes at the token we read (`</think>\n\n` = about to answer, `<think>\n`
# = about to reason). So |cos| scores refusal semantics *plus* mode, and on
# a heavily-merged base the mode term dominates. Robinson's stock
# Qwen3.8-27B scored 0.99 across that same split; this model scores 0.62,
# and that difference says more about the templates than the direction.
#
# Separation asks the question that actually predicts efficacy: does this
# direction split harmful from harmless prompts? Here it does, superbly
# (AUC 0.9996+ across the whole window) — the direction was never the
# problem, only where we removed it. Agreement is still reported, as a
# diagnostic rather than a selector.
# The sink screen is a FILTER on selection, not just a post-hoc abort.
# Separation and sink-energy both climb with depth on this model, so the
# best-separating layer (39, d=9.89) is also sink-dominated (1.97% > 1%)
# and would brick the model. Pick the best separator *among layers that
# pass the screen* — one pass, no guess-and-retry.
sink = {L: sink_energy(dirs[False][L]) for L in window}
by_sep = lambda L: sep[L][0]
eligible = [L for L in window if sink[L] <= SINK_ENERGY_MAX]
print(f"\nlayer selection over [{lo},{hi}] — separation, gated on sink < "
f"{SINK_ENERGY_MAX*100:.1f}%:")
for L in sorted(window, key=by_sep, reverse=True)[:8]:
mark = "ok " if sink[L] <= SINK_ENERGY_MAX else "SINK"
print(f" [{mark}] L{L:<3} d={sep[L][0]:6.3f} AUC={sep[L][1]:.4f} "
f"sink={sink[L]*100:6.3f}% |cos|={agree[L]:.4f}")
if not eligible:
print("\n!! every layer in the window is sink-dominated — no safe direction exists "
"here. Widen the window or reconsider the approach.", file=sys.stderr)
sys.exit(3)
best = max(eligible, key=by_sep)
layer = args.layer if args.layer is not None else best
d_unit = dirs[False][layer] # thinking-off direction at the chosen layer
top = sorted(window, key=lambda L: float(agree[L]), reverse=True)[:5]
print(f"\nagreement peak in [{lo},{hi}]: layer {best} (|cos|={float(agree[best]):.4f})")
print(" top-5 in window: " + ", ".join(f"L{L}={float(agree[L]):.4f}" for L in top))
print(f" global argmax (out-of-window layers are sink-dominated, informational only): "
f"L{int(agree.argmax())} (|cos|={float(agree.max()):.4f})")
print(f" using layer {layer} (|cos|={float(agree[layer]):.4f}); "
f"recipe anchor L{DEFAULT_LAYER} at |cos| 0.9925")
print(f" -> {len(eligible)}/{len(window)} layers pass the sink screen; "
f"best separator among them: L{best} (d={sep[best][0]:.3f})")
print(f" using layer {layer} (d={sep[layer][0]:.3f}, AUC={sep[layer][1]:.4f}, "
f"sink={sink[layer]*100:.3f}%, two-template |cos|={agree[layer]:.4f})")
agree_best = max(window, key=lambda L: agree[L])
print(f" [diagnostic] agreement would have picked L{agree_best} "
f"(|cos|={agree[agree_best]:.4f}, d={sep[agree_best][0]:.3f}) — "
f"recipe anchor L{DEFAULT_LAYER} at |cos| 0.9925 on stock Qwen3.8")
# --- finite gate: a NaN/Inf direction must NEVER pass silently -----------
# (the sink screen alone doesn't catch this — `nan > threshold` is False, so
@@ -428,8 +729,10 @@ def main():
blob = {
"layer": layer, "direction": d_unit.cpu(), "sink_energy": e,
"calibration": calib_prov,
"agreement": None if agree is None else float(agree[layer]),
"agreement_per_layer": None if agree is None else agree.cpu(),
"agreement": None if agree is None else agree[layer],
"agreement_per_layer": agree,
"separation": None if sep is None else sep[layer],
"separation_per_layer": sep,
"capture_window": CAPTURE_WINDOW,
}
@@ -439,47 +742,6 @@ def main():
print(f"direction saved -> {dpath} (capture-only, no write)")
return
if not args.out:
print("\n!! --out is required to write the abliterated model "
"(use --capture for direction-only).", file=sys.stderr)
sys.exit(1)
# --- surgery: orthogonalize every residual writer -------------------------
print(f"\northogonalizing {total_edits} residual writers along the refusal direction...")
sd = model.state_dict()
# Offload gate. `orthogonalize_` edits in place; a parameter that accelerate
# has offloaded shows up here as a meta tensor, where `sub_` writes into
# nothing and reports no error. That ships a quietly half-abliterated model —
# the same failure the coverage gate exists to prevent, arriving by a
# different door. Free the VRAM (stop the seats) rather than defeating this.
targets = trunk["down_proj"] + trunk["o_proj"] + trunk["linear_out"] + mtp_writers + embed
missing = [k for k in targets if k not in sd]
meta = [k for k in targets if k in sd and sd[k].device.type == "meta"]
if missing or meta:
print(f"\n!! surgery pre-check FAILED — {len(missing)} target tensor(s) absent from the "
f"state dict and {len(meta)} on the meta device (offloaded). In-place edits to "
f"those are silent no-ops. Ensure the model loads fully resident (stop the vLLM "
f"seats) and do not run --max-layer on the write path.", file=sys.stderr)
for k in (missing + meta)[:5]:
print(f" {k}", file=sys.stderr)
sys.exit(7)
print(f" -> surgery pre-check PASSED ({len(targets)} targets resident, none offloaded)")
edited = 0
for k in trunk["down_proj"] + trunk["o_proj"] + trunk["linear_out"] + mtp_writers:
orthogonalize_(sd[k], d_unit); edited += 1
for k in embed:
orthogonalize_embed_(sd[k], d_unit); edited += 1
print(f" edited {edited} tensors; vision ({n_vision}) untouched")
out_dir = Path(args.out); out_dir.mkdir(parents=True, exist_ok=True)
print(f"saving abliterated bf16 -> {out_dir}")
model.save_pretrained(out_dir, safe_serialization=True)
tok.save_pretrained(out_dir)
torch.save(blob, out_dir / "refusal-direction.pt")
print("DONE.")
if __name__ == "__main__":
main()