mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-29 15:31:25 -07:00
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:
co-authored by
kabachuha
joninco
Ashar
parent
ffa66af2d4
commit
ce8b05c77b
+55
-1
@@ -84,7 +84,7 @@ scorers = [
|
|||||||
# Note that only a single modifier can currently be applied,
|
# Note that only a single modifier can currently be applied,
|
||||||
# and this list must contain exactly one entry.
|
# and this list must contain exactly one entry.
|
||||||
modifiers = [
|
modifiers = [
|
||||||
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
|
{ plugin = "heretic.modifiers.ara.ARA" },
|
||||||
]
|
]
|
||||||
|
|
||||||
# Number of abliteration trials to run during optimization.
|
# Number of abliteration trials to run during optimization.
|
||||||
@@ -119,6 +119,7 @@ dataset = "mlabonne/harmful_behaviors"
|
|||||||
split = "train[:100]"
|
split = "train[:100]"
|
||||||
column = "text"
|
column = "text"
|
||||||
|
|
||||||
|
|
||||||
# Plugin-specific settings live in top-level TOML tables.
|
# 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).
|
# 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]"
|
split = "test[:100]"
|
||||||
column = "text"
|
column = "text"
|
||||||
|
|
||||||
|
|
||||||
# Dataset of prompts used to measure KL divergence from original model.
|
# Dataset of prompts used to measure KL divergence from original model.
|
||||||
[scorer.KLDivergence.prompts]
|
[scorer.KLDivergence.prompts]
|
||||||
dataset = "mlabonne/harmless_alpaca"
|
dataset = "mlabonne/harmless_alpaca"
|
||||||
split = "test[:100]"
|
split = "test[:100]"
|
||||||
column = "text"
|
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]
|
[modifier.Abliteration]
|
||||||
# Whether to adjust the residual directions so that only the component that is
|
# Whether to adjust the residual directions so that only the component that is
|
||||||
# orthogonal to the good direction is subtracted during abliteration.
|
# orthogonal to the good direction is subtracted during abliteration.
|
||||||
|
|||||||
@@ -392,7 +392,7 @@ class Settings(BaseSettings):
|
|||||||
modifiers: list[ModifierConfig] = Field(
|
modifiers: list[ModifierConfig] = Field(
|
||||||
default=[
|
default=[
|
||||||
ModifierConfig(
|
ModifierConfig(
|
||||||
plugin="heretic.modifiers.abliteration.Abliteration",
|
plugin="heretic.modifiers.ara.ARA",
|
||||||
),
|
),
|
||||||
],
|
],
|
||||||
description=(
|
description=(
|
||||||
|
|||||||
+2
-2
@@ -630,7 +630,7 @@ def run():
|
|||||||
print(f" * {name} = [bold]{value}[/]")
|
print(f" * {name} = [bold]{value}[/]")
|
||||||
print("* Resetting model...")
|
print("* Resetting model...")
|
||||||
modifier.reset_model(ctx)
|
modifier.reset_model(ctx)
|
||||||
print(f"* Modifying model ({modifier_name})...")
|
print(f"* Modifying model using {modifier_name}...")
|
||||||
modifier.modify_model(ctx, parameters)
|
modifier.modify_model(ctx, parameters)
|
||||||
print("* Evaluating...")
|
print("* Evaluating...")
|
||||||
scores = evaluator.get_scores()
|
scores = evaluator.get_scores()
|
||||||
@@ -868,7 +868,7 @@ def run():
|
|||||||
ctx = Context(settings=settings, model=model)
|
ctx = Context(settings=settings, model=model)
|
||||||
print("* Resetting model...")
|
print("* Resetting model...")
|
||||||
modifier.reset_model(ctx)
|
modifier.reset_model(ctx)
|
||||||
print(f"* Modifying model ({modifier_name})...")
|
print(f"* Modifying model using {modifier_name}...")
|
||||||
parameters = modifier.parameters_class.from_dict(
|
parameters = modifier.parameters_class.from_dict(
|
||||||
trial.user_attrs["parameters"]
|
trial.user_attrs["parameters"]
|
||||||
)
|
)
|
||||||
|
|||||||
+135
-1
@@ -2,12 +2,13 @@
|
|||||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||||
|
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from typing import Any, Type, cast
|
from typing import Any, Callable, Type, TypeAlias, cast
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from peft import LoraConfig, PeftModel, get_peft_model
|
from peft import LoraConfig, PeftModel, get_peft_model
|
||||||
from torch import FloatTensor, LongTensor, Tensor
|
from torch import FloatTensor, LongTensor, Tensor
|
||||||
from torch.nn import Module, ModuleList
|
from torch.nn import Module, ModuleList
|
||||||
|
from torch.utils.hooks import RemovableHandle
|
||||||
from transformers import (
|
from transformers import (
|
||||||
AutoModelForCausalLM,
|
AutoModelForCausalLM,
|
||||||
AutoModelForImageTextToText,
|
AutoModelForImageTextToText,
|
||||||
@@ -41,6 +42,13 @@ def get_model_class(
|
|||||||
return AutoModelForCausalLM
|
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:
|
class Model:
|
||||||
model: PreTrainedModel | PeftModel
|
model: PreTrainedModel | PeftModel
|
||||||
tokenizer: PreTrainedTokenizerBase
|
tokenizer: PreTrainedTokenizerBase
|
||||||
@@ -616,6 +624,132 @@ class Model:
|
|||||||
|
|
||||||
return (running_sum / total_count).to(torch.float32)
|
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:
|
def get_logits(self, prompts: list[Prompt]) -> Tensor:
|
||||||
# We only generate one token, and we return the raw logits over the vocabulary
|
# We only generate one token, and we return the raw logits over the vocabulary
|
||||||
# at that token position, for each prompt.
|
# at that token position, for each prompt.
|
||||||
|
|||||||
@@ -160,27 +160,27 @@ class Abliteration(Modifier[Parameters]):
|
|||||||
print(
|
print(
|
||||||
f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..."
|
f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..."
|
||||||
)
|
)
|
||||||
self.good_prompts = ctx.load_prompts(self.settings.good_prompts)
|
good_prompts = ctx.load_prompts(self.settings.good_prompts)
|
||||||
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
|
print(f"* [bold]{len(good_prompts)}[/] prompts loaded")
|
||||||
|
|
||||||
print()
|
print()
|
||||||
print(
|
print(
|
||||||
f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..."
|
f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..."
|
||||||
)
|
)
|
||||||
self.bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
|
bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
|
||||||
print(f"* [bold]{len(self.bad_prompts)}[/] prompts loaded")
|
print(f"* [bold]{len(bad_prompts)}[/] prompts loaded")
|
||||||
|
|
||||||
print()
|
print()
|
||||||
print("Calculating per-layer residual directions...")
|
print("Calculating per-layer residual directions...")
|
||||||
|
|
||||||
print("* Obtaining residual mean for good prompts...")
|
print("* Obtaining residual mean for good prompts...")
|
||||||
good_means = model.get_residuals_mean(
|
good_means = model.get_residuals_mean(
|
||||||
self.good_prompts,
|
good_prompts,
|
||||||
winsorization_quantile=self.settings.winsorization_quantile,
|
winsorization_quantile=self.settings.winsorization_quantile,
|
||||||
)
|
)
|
||||||
print("* Obtaining residual mean for bad prompts...")
|
print("* Obtaining residual mean for bad prompts...")
|
||||||
bad_means = model.get_residuals_mean(
|
bad_means = model.get_residuals_mean(
|
||||||
self.bad_prompts,
|
bad_prompts,
|
||||||
winsorization_quantile=self.settings.winsorization_quantile,
|
winsorization_quantile=self.settings.winsorization_quantile,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -9,6 +9,11 @@ print_debug_information = true
|
|||||||
|
|
||||||
batch_size = 2
|
batch_size = 2
|
||||||
max_response_length = 10
|
max_response_length = 10
|
||||||
|
|
||||||
|
modifiers = [
|
||||||
|
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
|
||||||
|
]
|
||||||
|
|
||||||
n_trials = 2
|
n_trials = 2
|
||||||
n_startup_trials = 1
|
n_startup_trials = 1
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,11 @@ print_debug_information = true
|
|||||||
|
|
||||||
batch_size = 2
|
batch_size = 2
|
||||||
max_response_length = 10
|
max_response_length = 10
|
||||||
|
|
||||||
|
modifiers = [
|
||||||
|
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
|
||||||
|
]
|
||||||
|
|
||||||
n_trials = 2
|
n_trials = 2
|
||||||
n_startup_trials = 1
|
n_startup_trials = 1
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,11 @@ print_debug_information = true
|
|||||||
|
|
||||||
batch_size = 2
|
batch_size = 2
|
||||||
max_response_length = 10
|
max_response_length = 10
|
||||||
|
|
||||||
|
modifiers = [
|
||||||
|
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
|
||||||
|
]
|
||||||
|
|
||||||
n_trials = 2
|
n_trials = 2
|
||||||
n_startup_trials = 1
|
n_startup_trials = 1
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,11 @@ print_debug_information = true
|
|||||||
|
|
||||||
batch_size = 2
|
batch_size = 2
|
||||||
max_response_length = 10
|
max_response_length = 10
|
||||||
|
|
||||||
|
modifiers = [
|
||||||
|
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
|
||||||
|
]
|
||||||
|
|
||||||
n_trials = 2
|
n_trials = 2
|
||||||
n_startup_trials = 1
|
n_startup_trials = 1
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,11 @@ print_debug_information = true
|
|||||||
|
|
||||||
batch_size = 2
|
batch_size = 2
|
||||||
max_response_length = 10
|
max_response_length = 10
|
||||||
|
|
||||||
|
modifiers = [
|
||||||
|
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
|
||||||
|
]
|
||||||
|
|
||||||
n_trials = 2
|
n_trials = 2
|
||||||
n_startup_trials = 1
|
n_startup_trials = 1
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user