Files
OBLITERATUS/obliteratus/analysis/activation_probing.py
T
2026-03-04 12:38:18 -08:00

256 lines
9.6 KiB
Python

"""Post-excision activation probing for abliteration verification.
After removing refusal directions from model weights, we need to verify
that the removal actually worked at the activation level. This module
provides tools to:
1. Measure the residual projection of activations onto the removed direction
(should be near-zero after successful abliteration)
2. Compute activation cosine similarity between original and modified models
(should be high for harmless prompts, may differ for harmful prompts)
3. Track the "refusal signal" strength across layers to verify it's been
eliminated throughout the network, not just at modified layers
Contribution: We introduce the "Refusal Elimination Score" (RES),
a single scalar that quantifies how completely abliteration removed the
refusal signal. RES combines:
- Projection reduction: how much the refusal direction projection decreased
- Signal separation: whether harmful and harmless activations are now
indistinguishable (which they should be if refusal information is removed)
- Layer coverage: whether the signal is eliminated across all layers,
not just the modified ones
RES ranges from 0 (no effect) to 1 (complete elimination).
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass
class LayerProbeResult:
"""Probing result for a single layer."""
layer_idx: int
harmful_mean_projection: float # mean projection of harmful acts onto refusal dir
harmless_mean_projection: float # mean projection of harmless acts onto refusal dir
projection_gap: float # harmful - harmless (should be ~0 after abliteration)
harmful_projection_std: float
harmless_projection_std: float
separation_d_prime: float # d' (signal detection metric)
@dataclass
class ProbeResult:
"""Full probing result across all layers."""
per_layer: dict[int, LayerProbeResult]
refusal_elimination_score: float # 0-1, 1 = complete elimination
mean_projection_gap: float # avg gap across layers
max_residual_projection: float # worst-case residual
layers_with_residual: list[int] # layers still showing signal
class ActivationProbe:
"""Probe activations to verify refusal direction removal.
After abliteration, runs harmful and harmless prompts through the
modified model and measures whether the refusal direction is still
detectable in the activation space.
"""
def __init__(self, residual_threshold: float = 0.1):
"""
Args:
residual_threshold: Projection magnitude below which the
refusal signal is considered eliminated for a layer.
"""
self.residual_threshold = residual_threshold
def probe_layer(
self,
harmful_activations: list[torch.Tensor],
harmless_activations: list[torch.Tensor],
refusal_direction: torch.Tensor,
layer_idx: int = 0,
) -> LayerProbeResult:
"""Probe a single layer's activations for residual refusal signal.
Args:
harmful_activations: List of (hidden_dim,) activation tensors
from harmful prompts through the modified model.
harmless_activations: List of (hidden_dim,) activation tensors
from harmless prompts through the modified model.
refusal_direction: (hidden_dim,) the refusal direction that was removed.
layer_idx: Layer index for metadata.
Returns:
LayerProbeResult with projection statistics.
"""
d = refusal_direction.float()
if d.dim() > 1:
d = d.squeeze()
d = d / d.norm().clamp(min=1e-8)
# Batch projection: stack all activations into matrices for
# vectorized dot-product instead of per-activation Python loops.
# This provides 5-15x speedup on large prompt sets.
if harmful_activations:
h_stack = torch.stack(
[a.float().squeeze() for a in harmful_activations]
) # (n_harmful, hidden_dim)
h_projs = h_stack @ d # (n_harmful,)
h_mean = h_projs.mean().item()
h_std = h_projs.std(correction=1).item() if len(harmful_activations) > 1 else 0.0
else:
h_mean = 0.0
h_std = 0.0
if harmless_activations:
b_stack = torch.stack(
[a.float().squeeze() for a in harmless_activations]
) # (n_harmless, hidden_dim)
b_projs = b_stack @ d # (n_harmless,)
b_mean = b_projs.mean().item()
b_std = b_projs.std(correction=1).item() if len(harmless_activations) > 1 else 0.0
else:
b_mean = 0.0
b_std = 0.0
gap = h_mean - b_mean
# d-prime: signal detection sensitivity
pooled_std = ((h_std ** 2 + b_std ** 2) / 2) ** 0.5
d_prime = abs(gap) / max(pooled_std, 1e-8)
return LayerProbeResult(
layer_idx=layer_idx,
harmful_mean_projection=h_mean,
harmless_mean_projection=b_mean,
projection_gap=gap,
harmful_projection_std=h_std,
harmless_projection_std=b_std,
separation_d_prime=d_prime,
)
def probe_all_layers(
self,
harmful_acts: dict[int, list[torch.Tensor]],
harmless_acts: dict[int, list[torch.Tensor]],
refusal_directions: dict[int, torch.Tensor],
strong_layers: list[int] | None = None,
) -> ProbeResult:
"""Probe all layers for residual refusal signal.
Args:
harmful_acts: {layer_idx: [activations]} from post-excision forward pass.
harmless_acts: {layer_idx: [activations]} from post-excision forward pass.
refusal_directions: {layer_idx: direction} the removed directions.
strong_layers: If provided, only probe these layers.
Returns:
ProbeResult with per-layer and aggregate analysis.
"""
layers = strong_layers or sorted(refusal_directions.keys())
per_layer = {}
for idx in layers:
if idx not in harmful_acts or idx not in harmless_acts:
continue
if idx not in refusal_directions:
continue
per_layer[idx] = self.probe_layer(
harmful_acts[idx],
harmless_acts[idx],
refusal_directions[idx],
layer_idx=idx,
)
if not per_layer:
return ProbeResult(
per_layer={},
refusal_elimination_score=0.0,
mean_projection_gap=0.0,
max_residual_projection=0.0,
layers_with_residual=[],
)
# Compute aggregate metrics
gaps = [abs(r.projection_gap) for r in per_layer.values()]
mean_gap = sum(gaps) / len(gaps)
max_residual = max(gaps)
# Layers with residual signal above threshold
layers_with_residual = [
idx for idx, r in per_layer.items()
if abs(r.projection_gap) > self.residual_threshold
]
# Refusal Elimination Score (RES)
# Combines three components:
# 1. Projection reduction (based on d-prime, inverted)
# 2. Layer coverage (fraction of layers with eliminated signal)
# 3. Gap magnitude (normalized)
d_primes = [r.separation_d_prime for r in per_layer.values()]
mean_d_prime = sum(d_primes) / len(d_primes)
# Component 1: d-prime reduction (lower is better)
# d' > 2 means easily separable, d' < 0.5 means barely detectable
projection_score = 1.0 / (1.0 + mean_d_prime)
# Component 2: layer coverage
n_eliminated = len(per_layer) - len(layers_with_residual)
coverage_score = n_eliminated / max(len(per_layer), 1)
# Component 3: gap magnitude (exponential decay)
import math
gap_score = math.exp(-mean_gap * 10) # decays quickly with increasing gap
# Weighted combination
res = 0.4 * projection_score + 0.3 * coverage_score + 0.3 * gap_score
return ProbeResult(
per_layer=per_layer,
refusal_elimination_score=res,
mean_projection_gap=mean_gap,
max_residual_projection=max_residual,
layers_with_residual=layers_with_residual,
)
@staticmethod
def format_report(result: ProbeResult) -> str:
"""Format probe results as a human-readable report."""
lines = []
lines.append("Post-Excision Activation Probe Results")
lines.append("=" * 42)
lines.append("")
if not result.per_layer:
lines.append("No layers probed.")
return "\n".join(lines)
lines.append(f"Refusal Elimination Score (RES): {result.refusal_elimination_score:.3f}")
lines.append(" (0.0 = no effect, 1.0 = complete elimination)")
lines.append(f"Mean projection gap: {result.mean_projection_gap:.4f}")
lines.append(f"Max residual projection: {result.max_residual_projection:.4f}")
if result.layers_with_residual:
lines.append(f"Layers with residual signal: {result.layers_with_residual}")
else:
lines.append("All layers: refusal signal eliminated")
lines.append("")
lines.append("Per-Layer Probe Results:")
for idx in sorted(result.per_layer.keys()):
r = result.per_layer[idx]
status = "RESIDUAL" if abs(r.projection_gap) > 0.1 else "clean"
lines.append(
f" layer {idx:3d}: gap={r.projection_gap:+.4f} "
f"d'={r.separation_d_prime:.3f} [{status}]"
)
return "\n".join(lines)