feat: re-implement ARA as a modifier plugin

Arbitrary-Rank Ablation (ARA) (Weidmann 2026) was originally introduced by @p-e-w in #211

Co-authored-by: kabachuha <artemkhrapov2001@yandex.ru>
Co-authored-by: joninco <joninco@bullpoint.org>
Co-authored-by: Ashar <coder3101@users.noreply.github.com>
This commit is contained in:
Philipp Emanuel Weidmann
2026-09-28 19:24:08 +05:30
co-authored by kabachuha joninco Ashar
parent ffa66af2d4
commit ce8b05c77b
11 changed files with 575 additions and 11 deletions
+55 -1
View File
@@ -84,7 +84,7 @@ scorers = [
# Note that only a single modifier can currently be applied,
# and this list must contain exactly one entry.
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
{ plugin = "heretic.modifiers.ara.ARA" },
]
# Number of abliteration trials to run during optimization.
@@ -119,6 +119,7 @@ dataset = "mlabonne/harmful_behaviors"
split = "train[:100]"
column = "text"
# Plugin-specific settings live in top-level TOML tables.
#
# For scorer plugins, use: `[scorer.<ClassName>]` (and optionally `[scorer.<ClassName>_<instance_name>]` for instance-related config).
@@ -198,12 +199,65 @@ dataset = "mlabonne/harmful_behaviors"
split = "test[:100]"
column = "text"
# Dataset of prompts used to measure KL divergence from original model.
[scorer.KLDivergence.prompts]
dataset = "mlabonne/harmless_alpaca"
split = "test[:100]"
column = "text"
[scorer.BenchmarkScore]
# Name that describes what the configured benchmark score measures.
score_name = "PIQA acc_norm"
# Task ID of the benchmark in the Language Model Evaluation Harness.
task = "piqa"
# Task metric to use as the benchmark score.
metric = "acc_norm,none"
[modifier.ARA]
# Whether to renormalize the rows of the modified matrices to preserve
# the original matrices' row magnitudes. This is believed to improve
# intelligence retention (see Lai 2025, "Magnitude-Preserving Orthogonal Ablation").
preserve_row_magnitudes = true
# The rank of the LoRA adapter to use.
# While mathematically, ARA is of "arbitrary" rank, experiments have shown that
# singular values tend to drop rapidly after a few dozen dimensions, and approximating
# the full transformation with a LoRA has many practical advantages.
lora_rank = 50
# Number of (outer) L-BFGS optimization steps to perform.
n_optimization_steps = 5
# Learning rate to use in the L-BFGS optimizer.
learning_rate = 1.0
# Maximum number of (inner) iterations to perform per (outer) L-BFGS optimization step.
max_iter = 20
# Number of past updates to store for approximating the Hessian matrix in the L-BFGS optimizer.
history_size = 10
# Whether to print the loss value for each L-BFGS optimization step.
print_loss = false
# Dataset of prompts that tend to produce desirable responses.
[modifier.ARA.good_prompts]
dataset = "mlabonne/harmless_alpaca"
split = "train[:400]"
column = "text"
# Dataset of prompts that tend to produce undesirable responses.
[modifier.ARA.bad_prompts]
dataset = "mlabonne/harmful_behaviors"
split = "train[:400]"
column = "text"
[modifier.Abliteration]
# Whether to adjust the residual directions so that only the component that is
# orthogonal to the good direction is subtracted during abliteration.
+1 -1
View File
@@ -392,7 +392,7 @@ class Settings(BaseSettings):
modifiers: list[ModifierConfig] = Field(
default=[
ModifierConfig(
plugin="heretic.modifiers.abliteration.Abliteration",
plugin="heretic.modifiers.ara.ARA",
),
],
description=(
+2 -2
View File
@@ -630,7 +630,7 @@ def run():
print(f" * {name} = [bold]{value}[/]")
print("* Resetting model...")
modifier.reset_model(ctx)
print(f"* Modifying model ({modifier_name})...")
print(f"* Modifying model using {modifier_name}...")
modifier.modify_model(ctx, parameters)
print("* Evaluating...")
scores = evaluator.get_scores()
@@ -868,7 +868,7 @@ def run():
ctx = Context(settings=settings, model=model)
print("* Resetting model...")
modifier.reset_model(ctx)
print(f"* Modifying model ({modifier_name})...")
print(f"* Modifying model using {modifier_name}...")
parameters = modifier.parameters_class.from_dict(
trial.user_attrs["parameters"]
)
+135 -1
View File
@@ -2,12 +2,13 @@
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
from contextlib import suppress
from typing import Any, Type, cast
from typing import Any, Callable, Type, TypeAlias, cast
import torch
from peft import LoraConfig, PeftModel, get_peft_model
from torch import FloatTensor, LongTensor, Tensor
from torch.nn import Module, ModuleList
from torch.utils.hooks import RemovableHandle
from transformers import (
AutoModelForCausalLM,
AutoModelForImageTextToText,
@@ -41,6 +42,13 @@ def get_model_class(
return AutoModelForCausalLM
# The list contains one element per layer.
# Each element maps from the component name to a (possibly sparse) mapping
# from the module index to an (input, output) tuple containing the I/O
# tensors of shape (prompt, component).
ModuleIO: TypeAlias = list[dict[str, dict[int, tuple[Tensor, Tensor]]]]
class Model:
model: PreTrainedModel | PeftModel
tokenizer: PreTrainedTokenizerBase
@@ -616,6 +624,132 @@ class Model:
return (running_sum / total_count).to(torch.float32)
def get_module_io(
self,
prompts: list[Prompt],
) -> ModuleIO:
# The list contains one element per layer.
# Each element maps from the component name to a (possibly sparse) mapping
# from the module index to an (input, output) tuple containing the I/O
# tensors of shape (prompt, component).
module_io: ModuleIO = []
def get_hook(
layer_index: int,
component: str,
module_index: int,
) -> Callable[[Module, tuple[Tensor, ...], Tensor], None]:
def hook(
module: Module,
inputs: tuple[Tensor, ...],
outputs: Tensor,
) -> None:
if len(module_io) == layer_index:
# First invocation of the hook for this layer.
module_io.append({})
# Layers are invoked in order during inference,
# so this should always hold.
assert len(module_io) == layer_index + 1
if component not in module_io[layer_index]:
module_io[layer_index][component] = {}
# Each module should be invoked at most once per inference step.
assert module_index not in module_io[layer_index][component]
# inputs[0] and outputs have shape (prompt, position, component),
# so this extracts the input/output at the end of each prompt.
# Move to CPU to decouple from device assignments, which can
# change between model reloads in multi-GPU configurations.
input = inputs[0][:, -1, :].detach().clone().cpu()
output = outputs[:, -1, :].detach().clone().cpu()
# The modules associated with a component (e.g. expert MLPs)
# are not necessarily invoked in order, nor are all of them
# necessarily invoked in each inference step, so we cannot
# use a list here.
module_io[layer_index][component][module_index] = (input, output)
return hook
hook_handles: list[RemovableHandle] = []
for layer_index in range(len(self.get_layers())):
for component, modules in self.get_layer_modules(layer_index).items():
for module_index, module in enumerate(modules):
hook_handles.append(
module.register_forward_hook(
get_hook(layer_index, component, module_index)
)
)
self.generate(prompts, max_new_tokens=1)
for hook_handle in hook_handles:
hook_handle.remove()
return module_io
def get_module_io_batched(
self,
prompts: list[Prompt],
) -> ModuleIO:
# Aggregating batch results is more complicated for module I/O
# than for other get_*_batched methods, because the structure of the results
# might differ between batches, as whether individual modules activate
# can depend on the prompt (in particular for MoE models).
# In practice, inhomogeneous results should be very rare, but to be fully
# generic, this logic is required.
module_io_batches: list[ModuleIO] = [
self.get_module_io(batch)
for batch in batchify(prompts, self.settings.batch_size)
]
module_io: ModuleIO = []
for layer_index in range(len(self.get_layers())):
module_io.append({})
for module_io_batch in module_io_batches:
for component, io_map in module_io_batch[layer_index].items():
if component not in module_io[layer_index]:
module_io[layer_index][component] = {}
for module_index in io_map:
if module_index not in module_io[layer_index][component]:
# This is a placeholder; the actual aggregation happens below.
# We need to iterate over the batches twice because we don't
# know in advance which components and module indices are present.
module_io[layer_index][component][module_index] = (
torch.empty(0),
torch.empty(0),
)
for component, io_map in module_io[layer_index].items():
for module_index in io_map:
inputs_outputs = [
module_io_batch[layer_index][component][module_index]
for module_io_batch in module_io_batches
if component in module_io_batch[layer_index]
and module_index in module_io_batch[layer_index][component]
]
input = torch.cat(
[input_output[0] for input_output in inputs_outputs],
dim=0,
)
output = torch.cat(
[input_output[1] for input_output in inputs_outputs],
dim=0,
)
# The key already exists, and replacing existing values
# in a dictionary while iterating over the same dictionary
# is safe in Python.
module_io[layer_index][component][module_index] = (input, output)
return module_io
def get_logits(self, prompts: list[Prompt]) -> Tensor:
# We only generate one token, and we return the raw logits over the vocabulary
# at that token position, for each prompt.
+6 -6
View File
@@ -160,27 +160,27 @@ class Abliteration(Modifier[Parameters]):
print(
f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..."
)
self.good_prompts = ctx.load_prompts(self.settings.good_prompts)
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
good_prompts = ctx.load_prompts(self.settings.good_prompts)
print(f"* [bold]{len(good_prompts)}[/] prompts loaded")
print()
print(
f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..."
)
self.bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
print(f"* [bold]{len(self.bad_prompts)}[/] prompts loaded")
bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
print(f"* [bold]{len(bad_prompts)}[/] prompts loaded")
print()
print("Calculating per-layer residual directions...")
print("* Obtaining residual mean for good prompts...")
good_means = model.get_residuals_mean(
self.good_prompts,
good_prompts,
winsorization_quantile=self.settings.winsorization_quantile,
)
print("* Obtaining residual mean for bad prompts...")
bad_means = model.get_residuals_mean(
self.bad_prompts,
bad_prompts,
winsorization_quantile=self.settings.winsorization_quantile,
)
+351
View File
@@ -0,0 +1,351 @@
# SPDX-License-Identifier: AGPL-3.0-or-later
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
# Arbitrary-Rank Ablation (ARA) (Weidmann 2026)
# See https://github.com/p-e-w/heretic/pull/211 for more information.
from dataclasses import asdict, dataclass
from typing import Any, cast
import bitsandbytes as bnb
import torch
import torch.linalg as LA
import torch.nn.functional as F
from optuna import Trial
from peft.tuners.lora.layer import Linear
from pydantic import (
BaseModel,
Field,
PositiveInt,
)
from torch import Tensor
from torch.optim import LBFGS
from heretic.config import DatasetSpecification, SingleDatasetSpecification
from heretic.modifier import Context, Modifier, Serializable
from heretic.utils import format_dataset_specification, print
@dataclass
class Parameters(Serializable):
start_layer_index: int
end_layer_index: int
preserve_good_behavior_weight: float
steer_bad_behavior_weight: float
overcorrect_relative_weight: float
neighbor_count: int
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def to_presentation_dict(self) -> dict[str, str]:
return {
name: (f"{value:.4f}" if isinstance(value, float) else f"{value}")
for name, value in asdict(self).items()
}
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "Serializable":
return Parameters(**data)
class Settings(BaseModel):
good_prompts: DatasetSpecification = Field(
default=SingleDatasetSpecification(
dataset="mlabonne/harmless_alpaca",
split="train[:400]",
column="text",
),
description="Dataset of prompts that tend to produce desirable responses.",
)
bad_prompts: DatasetSpecification = Field(
default=SingleDatasetSpecification(
dataset="mlabonne/harmful_behaviors",
split="train[:400]",
column="text",
),
description="Dataset of prompts that tend to produce undesirable responses.",
)
preserve_row_magnitudes: bool = Field(
default=True,
description=(
"Whether to renormalize the rows of the modified matrices to preserve "
"the original matrices' row magnitudes. This is believed to improve "
'intelligence retention (see Lai 2025, "Magnitude-Preserving Orthogonal Ablation").'
),
)
lora_rank: PositiveInt = Field(
default=50,
description=(
"The rank of the LoRA adapter to use. "
'While mathematically, ARA is of "arbitrary" rank, experiments have shown that '
"singular values tend to drop rapidly after a few dozen dimensions, and approximating "
"the full transformation with a LoRA has many practical advantages."
),
)
n_optimization_steps: PositiveInt = Field(
default=5,
description="Number of (outer) L-BFGS optimization steps to perform.",
)
learning_rate: float = Field(
default=1.0,
description="Learning rate to use in the L-BFGS optimizer.",
)
max_iter: PositiveInt = Field(
default=20,
description="Maximum number of (inner) iterations to perform per (outer) L-BFGS optimization step.",
)
history_size: PositiveInt = Field(
default=10,
description="Number of past updates to store for approximating the Hessian matrix in the L-BFGS optimizer.",
)
print_loss: bool = Field(
default=False,
description="Whether to print the loss value for each L-BFGS optimization step.",
)
# For each vector in the 2D-tensor `a`, computes the mean Euclidean distance
# to the `k` nearest neighbors of the vector among the vectors in the 2D-tensor `b`.
def mean_distances_to_knn(a: Tensor, b: Tensor, k: int) -> Tensor:
distances = torch.cdist(a, b)
nearest_distances, _ = distances.topk(k, dim=1, largest=False)
return nearest_distances.mean(1)
# The objective function at the heart of ARA.
def ara_loss(
good_output: Tensor,
bad_output: Tensor,
new_good_output: Tensor,
new_bad_output: Tensor,
parameters: Parameters,
) -> Tensor:
# The outputs for "good" prompts should change as little as possible.
preserve_good_behavior = ((new_good_output - good_output) ** 2).mean()
steer_bad_behavior = (
# Pull the outputs for "bad" prompts towards
# the original outputs for "good" prompts.
mean_distances_to_knn(
new_bad_output,
good_output,
parameters.neighbor_count,
).mean()
# Push the outputs for "bad" prompts away from
# the original outputs for "bad" prompts.
# In combination with the above, this overcorrects
# away from the original residuals, which results
# in stronger steering that can overcome more complex
# refusal mechanisms.
+ parameters.overcorrect_relative_weight
* -mean_distances_to_knn(
new_bad_output,
bad_output,
parameters.neighbor_count,
).mean()
)
return (
parameters.preserve_good_behavior_weight * preserve_good_behavior
+ parameters.steer_bad_behavior_weight * steer_bad_behavior
)
class ARA(Modifier[Parameters]):
settings: Settings
@property
def reproducible(self) -> bool:
return True
@property
def modifier_name(self) -> str:
return "Arbitrary-Rank Ablation (ARA)"
def init(self, ctx: Context) -> None:
model = ctx.get_model()
print()
print(
f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..."
)
good_prompts = ctx.load_prompts(self.settings.good_prompts)
print(f"* [bold]{len(good_prompts)}[/] prompts loaded")
print()
print(
f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..."
)
bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
print(f"* [bold]{len(bad_prompts)}[/] prompts loaded")
print()
print("Obtaining module I/O for good prompts...")
self.good_module_io = model.get_module_io_batched(good_prompts)
print()
print("Obtaining module I/O for bad prompts...")
self.bad_module_io = model.get_module_io_batched(bad_prompts)
# LoRA B matrices are initialized to zero by default in PEFT,
# so we don't need to do anything manually.
model.apply_lora(self.settings.lora_rank)
def suggest_parameters(self, ctx: Context, trial: Trial) -> Parameters:
layer_count = len(ctx.get_model().get_layers())
start_layer_index = trial.suggest_int(
"start_layer_index",
0,
layer_count // 2,
)
end_layer_index = trial.suggest_int(
"end_layer_index",
layer_count // 2,
layer_count,
)
preserve_good_behavior_weight = trial.suggest_float(
"preserve_good_behavior_weight",
0.0,
1.0,
)
steer_bad_behavior_weight = trial.suggest_float(
"steer_bad_behavior_weight",
0.0001,
1.0,
log=True,
)
overcorrect_relative_weight = trial.suggest_float(
"overcorrect_relative_weight",
0.0,
1.3,
)
neighbor_count = trial.suggest_int(
"neighbor_count",
1,
15,
)
return Parameters(
start_layer_index=start_layer_index,
end_layer_index=end_layer_index,
preserve_good_behavior_weight=preserve_good_behavior_weight,
steer_bad_behavior_weight=steer_bad_behavior_weight,
overcorrect_relative_weight=overcorrect_relative_weight,
neighbor_count=neighbor_count,
)
def modify_model(self, ctx: Context, parameters: Parameters) -> None:
model = ctx.get_model()
for layer_index in range(
parameters.start_layer_index,
parameters.end_layer_index,
):
for component, modules in model.get_layer_modules(layer_index).items():
for module_index, module in enumerate(modules):
# Cast to Linear to access weights and LoRA adapters.
module = cast(Linear, module)
# We need the base weight in float32 to compute the effective weight.
base_weight = cast(Tensor, module.base_layer.weight)
quant_state = getattr(base_weight, "quant_state", None)
if quant_state is None:
W_base = base_weight.to(torch.float32)
else:
# Use the original dequantization logic from bitsandbytes.
W_base = cast(
Tensor,
bnb.functional.dequantize_4bit( # ty:ignore[possibly-missing-attribute]
base_weight.data,
quant_state,
).to(torch.float32),
)
# Pre-calculate the original row norms to preserve them.
# See https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration
W_row_norms = cast(
Tensor,
LA.vector_norm(W_base, dim=1, keepdim=True).detach(),
)
# We optimize the LoRA weights A and B.
lora_A = cast(Tensor, module.lora_A["default"].weight)
lora_B = cast(Tensor, module.lora_B["default"].weight)
# Move I/O tensors to the device of the adapter weights.
good_input, good_output = self.good_module_io[layer_index][
component
][module_index]
bad_input, bad_output = self.bad_module_io[layer_index][component][
module_index
]
good_input = good_input.float().to(lora_A.device)
good_output = good_output.float().to(lora_A.device)
bad_input = bad_input.float().to(lora_A.device)
bad_output = bad_output.float().to(lora_A.device)
def objective(A: Tensor, B: Tensor) -> Tensor:
# Calculate effective weight after applying adapter.
W_eff = W_base + (B @ A)
if self.settings.preserve_row_magnitudes:
# Normalize to unit length, then scale by original norms,
# preserving the original row norms.
W_eff = F.normalize(W_eff, p=2, dim=1) * W_row_norms
# Compute outputs using the effective weight.
new_good_output = good_input @ W_eff.T
new_bad_output = bad_input @ W_eff.T
return ara_loss(
good_output,
bad_output,
new_good_output,
new_bad_output,
parameters,
)
optimizer = LBFGS(
[lora_A, lora_B],
lr=self.settings.learning_rate,
max_iter=self.settings.max_iter,
history_size=self.settings.history_size,
line_search_fn="strong_wolfe",
)
def closure() -> Tensor:
optimizer.zero_grad()
loss = objective(lora_A, lora_B)
loss.backward()
return loss
for step in range(self.settings.n_optimization_steps):
loss = optimizer.step(closure)
if self.settings.print_loss:
print(
f"\\[{layer_index}/{component}/{module_index}] Step: {step + 1}, Loss: {loss.item():.6f}"
)
# Free the gradient buffers accumulated during optimization.
# Without this, they persist on the model (one full-size gradient
# per processed weight) and can easily consume tens of GB of VRAM,
# causing out-of-memory errors during the subsequent evaluation.
optimizer.zero_grad(set_to_none=True)
def reset_model(self, ctx: Context) -> None:
model = ctx.get_model()
fast_path = model.reset_model()
if not fast_path:
model.apply_lora(self.settings.lora_rank)
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2
max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2
n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2
max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2
n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2
max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2
n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2
max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2
n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2
max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2
n_startup_trials = 1