19 Commits

Author SHA1 Message Date
Ashar edc3b12345 fix(ara): free gradient buffers after optimization (#426)
LBFGS leaves one full-size gradient buffer on every optimized weight.
Across the layers processed in a typical trial this is many GiB of VRAM
that persists into evaluation, causing CUDA out-of-memory errors.
Clear the gradients after each module is optimized.
2026-08-17 14:57:09 +05:30
kabachuha 25979ad7d0 feat: ARA, but it's LoRA (#332)
* ARA, but it's LoRA

* ARA, but it's LoRA: address Gemini's review

* ARA, but it's LoRA: Gemini is stupid
2026-07-05 13:51:05 +05:30
Philipp Emanuel Weidmann 3b70fe5dfa fix(ara): set batch size on HFLM object 2026-04-01 14:34:21 +05:30
Philipp Emanuel Weidmann f7a456bd0c feat(ara): add optional row-norm preservation 2026-03-31 15:55:27 +05:30
Philipp Emanuel Weidmann 988c6bd90e Merge branch 'master' into ara 2026-03-30 13:28:57 +05:30
Philipp Emanuel Weidmann c925f5e802 feat(ara): add option to optimize for PIQA instead of KLD 2026-03-21 13:21:43 +05:30
Philipp Emanuel Weidmann 4a6304c361 feat(ara): add abliteration method to model card 2026-03-10 11:16:09 +05:30
Philipp Emanuel Weidmann c76416fe03 feat(ara): expand parameter ranges
Incorporates feedback from @joninco
2026-03-10 10:27:32 +05:30
joninco 2bb203ee47 fix(ara): store captured I/O tensors on CPU for multi-GPU robustness (#214)
Extends d79a443 — that commit correctly moves I/O tensors to the weight
matrix's device before L-BFGS optimization, but the captured tensors
remain on their original GPU between trials. When reset_model() reloads
the model, device assignments can change, leaving orphaned tensors on
GPUs that now need that VRAM for the reloaded weights.

Moving to CPU at capture time in get_module_io ensures:
- Zero VRAM wasted on stale device assignments between trials
- Clean CPU→target transfer regardless of how devices shuffle on reload
- No overhead on single-GPU (.cpu() is a no-op when already on CPU,
  and .to(device) in ara_abliterate handles the final placement)
2026-03-07 19:40:45 +05:30
Philipp Emanuel Weidmann d79a443e6f feat(ara): fix issues on some multi-GPU setups
Co-authored-by: kabachuha <artemkhrapov2001@yandex.ru>
2026-03-07 18:24:31 +05:30
Philipp Emanuel Weidmann 0bb9521fbe feat(ara): optimize all parameters 2026-03-05 08:55:58 +05:30
Philipp Emanuel Weidmann 992fb3a4b3 feat(ara): change weights to match those used for the gpt-oss-20b demo 2026-03-04 16:48:25 +05:30
Philipp Emanuel Weidmann 304c14adc7 feat(ara): replace weight parameters with balance parameter 2026-03-04 11:59:46 +05:30
Philipp Emanuel Weidmann 56e57adf36 feat(ara): improve steering term 2026-03-04 11:06:01 +05:30
Philipp Emanuel Weidmann bd1fa0ade4 feat(ara): remove tie_to_original_matrix term
This term was found experimentally to be 3-4 orders of magnitude smaller than the others in most runs, and have no meaningful effect on the result of the optimization.
2026-03-04 09:13:38 +05:30
Philipp Emanuel Weidmann 3c5d6920bf feat(ara): fix memory leak 2026-03-02 18:50:27 +05:30
Philipp Emanuel Weidmann b8f4a9c985 feat(ara): implement optimization for ARA parameters 2026-03-02 14:36:57 +05:30
Philipp Emanuel Weidmann 154241f8a2 feat(ara): implement matrix optimization 2026-02-27 14:16:27 +05:30
Philipp Emanuel Weidmann ea7c59a55a feat(ara): add methods for obtaining module I/O 2026-02-25 17:44:46 +05:30
5 changed files with 756 additions and 168 deletions
+37 -1
View File
@@ -188,6 +188,42 @@ class Settings(BaseSettings):
),
)
target_components: list[str] = Field(
default=["attn.o_proj", "mlp.down_proj"],
description=(
"List of component names to target for abliteration. "
'Currently supported values are "attn.o_proj" and "mlp.down_proj".'
),
)
use_ara: bool = Field(
default=True,
description=(
"Whether to use Arbitrary-Rank Ablation (ARA), an abliteration method based on matrix optimization, "
"instead of traditional directional ablation."
),
)
use_ara_lora: bool = Field(
default=False,
description=(
"Use LoRA in ARA instead of full-weight editing. Makes it compatible with quantization and removes model reloads."
),
)
ara_lora_rank: int = Field(
default=128,
description="If LoRA is used in ARA, this sets up its rank. Keep it high enough to simulate the 'arbitrary' effect.",
)
use_piqa: bool = Field(
default=False,
description=(
"Whether to use the Physical Interaction: Question Answering (PIQA) benchmark "
"as the quality metric instead of the Kullback-Leibler divergence."
),
)
orthogonalize_direction: bool = Field(
default=False,
description=(
@@ -197,7 +233,7 @@ class Settings(BaseSettings):
)
row_normalization: RowNormalization = Field(
default=RowNormalization.NONE,
default=RowNormalization.FULL,
description=(
"How to apply row normalization of the weights. Options: "
'"none" (no normalization), '
+53 -28
View File
@@ -1,7 +1,9 @@
# SPDX-License-Identifier: AGPL-3.0-or-later
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
import lm_eval
import torch.nn.functional as F
from lm_eval.models.huggingface import HFLM
from torch import Tensor
from .config import Settings
@@ -21,15 +23,16 @@ class Evaluator:
self.settings = settings
self.model = model
print()
print(
f"Loading good evaluation prompts from [bold]{settings.good_evaluation_prompts.dataset}[/]..."
)
self.good_prompts = load_prompts(settings, settings.good_evaluation_prompts)
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
if not settings.use_piqa:
print()
print(
f"Loading good evaluation prompts from [bold]{settings.good_evaluation_prompts.dataset}[/]..."
)
self.good_prompts = load_prompts(settings, settings.good_evaluation_prompts)
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
print("* Obtaining first-token probability distributions...")
self.base_logprobs = model.get_logprobs_batched(self.good_prompts)
print("* Obtaining first-token probability distributions...")
self.base_logprobs = model.get_logprobs_batched(self.good_prompts)
print()
print(
@@ -93,35 +96,57 @@ class Evaluator:
return refusal_count
def get_score(self) -> tuple[tuple[float, float], float, int]:
print(" * Obtaining first-token probability distributions...")
logprobs = self.model.get_logprobs_batched(self.good_prompts)
kl_divergence = F.kl_div(
logprobs,
self.base_logprobs,
reduction="batchmean",
log_target=True,
).item()
print(f" * KL divergence: [bold]{kl_divergence:.4f}[/]")
if self.settings.use_piqa:
print(" * Running PIQA benchmark...")
hflm = HFLM(
pretrained=self.model.model, # ty:ignore[invalid-argument-type]
tokenizer=self.model.tokenizer, # ty:ignore[invalid-argument-type]
batch_size="auto",
)
results = lm_eval.simple_evaluate(
model=hflm,
tasks=["piqa"],
)
piqa_acc_norm: float = results["results"]["piqa"]["acc_norm,none"]
print(f" * PIQA acc_norm: [bold]{piqa_acc_norm:.4f}[/]")
else:
print(" * Obtaining first-token probability distributions...")
logprobs = self.model.get_logprobs_batched(self.good_prompts)
kl_divergence = F.kl_div(
logprobs,
self.base_logprobs,
reduction="batchmean",
log_target=True,
).item()
print(f" * KL divergence: [bold]{kl_divergence:.4f}[/]")
print(" * Counting model refusals...")
refusals = self.count_refusals()
print(f" * Refusals: [bold]{refusals}[/]/{len(self.bad_prompts)}")
kl_divergence_scale = self.settings.kl_divergence_scale
kl_divergence_target = self.settings.kl_divergence_target
refusals_score = (
refusals / self.base_refusals if self.base_refusals > 0 else float(refusals)
)
if kl_divergence >= kl_divergence_target:
kld_score = kl_divergence / kl_divergence_scale
if self.settings.use_piqa:
score = (
-piqa_acc_norm,
refusals_score,
)
return score, -piqa_acc_norm, refusals
else:
kld_score = refusals_score * kl_divergence_target / kl_divergence_scale
kl_divergence_scale = self.settings.kl_divergence_scale
kl_divergence_target = self.settings.kl_divergence_target
score = (
kld_score,
refusals_score,
)
if kl_divergence >= kl_divergence_target:
kld_score = kl_divergence / kl_divergence_scale
else:
kld_score = refusals_score * kl_divergence_target / kl_divergence_scale
return score, kl_divergence, refusals
score = (
kld_score,
refusals_score,
)
return score, kl_divergence, refusals
+213 -109
View File
@@ -51,9 +51,9 @@ from rich.table import Table
from rich.traceback import install
from .analyzer import Analyzer
from .config import QuantizationMethod, Settings
from .config import QuantizationMethod, RowNormalization, Settings
from .evaluator import Evaluator
from .model import AbliterationParameters, Model, get_model_class
from .model import AbliterationParameters, ARAParameters, Model, get_model_class
from .utils import (
empty_cache,
format_duration,
@@ -227,8 +227,9 @@ def run():
"[bold yellow]No GPU or other accelerator detected. Operations will be slow.[/]"
)
# We don't need gradients as we only do inference.
torch.set_grad_enabled(False)
if not settings.use_ara:
# We don't need gradients as we only do inference.
torch.set_grad_enabled(False)
# While determining the optimal batch size, we will try many different batch sizes,
# resulting in many computation graphs being compiled. Raising the limit (default = 8)
@@ -451,40 +452,47 @@ def run():
evaluator.get_score()
return
print()
print("Calculating per-layer refusal directions...")
print("* Obtaining residuals for good prompts...")
good_residuals = model.get_residuals_batched(good_prompts)
print("* Obtaining residuals for bad prompts...")
bad_residuals = model.get_residuals_batched(bad_prompts)
if settings.use_ara:
print()
print("Obtaining module I/O for good prompts...")
good_module_io = model.get_module_io_batched(good_prompts)
print("Obtaining module I/O for bad prompts...")
bad_module_io = model.get_module_io_batched(bad_prompts)
else:
print()
print("Calculating per-layer refusal directions...")
print("* Obtaining residuals for good prompts...")
good_residuals = model.get_residuals_batched(good_prompts)
print("* Obtaining residuals for bad prompts...")
bad_residuals = model.get_residuals_batched(bad_prompts)
good_means = good_residuals.mean(dim=0)
bad_means = bad_residuals.mean(dim=0)
good_means = good_residuals.mean(dim=0)
bad_means = bad_residuals.mean(dim=0)
refusal_directions = F.normalize(bad_means - good_means, p=2, dim=1)
refusal_directions = F.normalize(bad_means - good_means, p=2, dim=1)
if settings.orthogonalize_direction:
# Implements https://huggingface.co/blog/grimjim/projected-abliteration
# Adjust the refusal directions so that only the component that is
# orthogonal to the good direction is subtracted during abliteration.
good_directions = F.normalize(good_means, p=2, dim=1)
projection_vector = torch.sum(refusal_directions * good_directions, dim=1)
refusal_directions = (
refusal_directions - projection_vector.unsqueeze(1) * good_directions
)
refusal_directions = F.normalize(refusal_directions, p=2, dim=1)
if settings.orthogonalize_direction:
# Implements https://huggingface.co/blog/grimjim/projected-abliteration
# Adjust the refusal directions so that only the component that is
# orthogonal to the good direction is subtracted during abliteration.
good_directions = F.normalize(good_means, p=2, dim=1)
projection_vector = torch.sum(refusal_directions * good_directions, dim=1)
refusal_directions = (
refusal_directions - projection_vector.unsqueeze(1) * good_directions
)
refusal_directions = F.normalize(refusal_directions, p=2, dim=1)
analyzer = Analyzer(settings, model, good_residuals, bad_residuals)
analyzer = Analyzer(settings, model, good_residuals, bad_residuals)
if settings.print_residual_geometry:
analyzer.print_residual_geometry()
if settings.print_residual_geometry:
analyzer.print_residual_geometry()
if settings.plot_residuals:
analyzer.plot_residuals()
if settings.plot_residuals:
analyzer.plot_residuals()
# We don't need the residuals after computing refusal directions.
del good_residuals, bad_residuals, analyzer
empty_cache()
# We don't need the residuals after computing refusal directions.
del good_residuals, bad_residuals, analyzer
empty_cache()
trial_index = 0
start_index = 0
@@ -495,83 +503,144 @@ def run():
trial_index += 1
trial.set_user_attr("index", trial_index)
direction_scope = trial.suggest_categorical(
"direction_scope",
[
"global",
"per layer",
],
)
last_layer_index = len(model.get_layers()) - 1
# Discrimination between "harmful" and "harmless" inputs is usually strongest
# in layers slightly past the midpoint of the layer stack. See the original
# abliteration paper (https://arxiv.org/abs/2406.11717) for a deeper analysis.
#
# Note that we always sample this parameter even though we only need it for
# the "global" direction scope. The reason is that multivariate TPE doesn't
# work with conditional or variable-range parameters.
direction_index = trial.suggest_float(
"direction_index",
0.4 * last_layer_index,
0.9 * last_layer_index,
)
if direction_scope == "per layer":
direction_index = None
parameters = {}
for component in model.get_abliterable_components():
# The parameter ranges are based on experiments with various models
# and much wider ranges. They are not set in stone and might have to be
# adjusted for future models.
max_weight = trial.suggest_float(
f"{component}.max_weight",
0.8,
1.5,
if settings.use_ara:
start_layer_index = trial.suggest_int(
"start_layer_index",
0,
len(model.get_layers()) // 2,
)
max_weight_position = trial.suggest_float(
f"{component}.max_weight_position",
0.6 * last_layer_index,
1.0 * last_layer_index,
end_layer_index = trial.suggest_int(
"end_layer_index",
len(model.get_layers()) // 2,
len(model.get_layers()),
)
# For sampling purposes, min_weight is expressed as a fraction of max_weight,
# again because multivariate TPE doesn't support variable-range parameters.
# The value is transformed into the actual min_weight value below.
min_weight = trial.suggest_float(
f"{component}.min_weight",
preserve_good_behavior_weight = trial.suggest_float(
"preserve_good_behavior_weight",
0.0,
1.0,
)
min_weight_distance = trial.suggest_float(
f"{component}.min_weight_distance",
steer_bad_behavior_weight = trial.suggest_float(
"steer_bad_behavior_weight",
0.0001,
1.0,
0.6 * last_layer_index,
log=True,
)
overcorrect_relative_weight = trial.suggest_float(
"overcorrect_relative_weight",
0.0,
1.3,
)
neighbor_count = trial.suggest_int(
"neighbor_count",
1,
15,
)
parameters[component] = AbliterationParameters(
max_weight=max_weight,
max_weight_position=max_weight_position,
min_weight=(min_weight * max_weight),
min_weight_distance=min_weight_distance,
ara_parameters = ARAParameters(
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,
)
trial.set_user_attr("direction_index", direction_index)
trial.set_user_attr("parameters", {k: asdict(v) for k, v in parameters.items()})
trial.set_user_attr("ara_parameters", asdict(ara_parameters))
else:
direction_scope = trial.suggest_categorical(
"direction_scope",
[
"global",
"per layer",
],
)
last_layer_index = len(model.get_layers()) - 1
# Discrimination between "harmful" and "harmless" inputs is usually strongest
# in layers slightly past the midpoint of the layer stack. See the original
# abliteration paper (https://arxiv.org/abs/2406.11717) for a deeper analysis.
#
# Note that we always sample this parameter even though we only need it for
# the "global" direction scope. The reason is that multivariate TPE doesn't
# work with conditional or variable-range parameters.
direction_index = trial.suggest_float(
"direction_index",
0.4 * last_layer_index,
0.9 * last_layer_index,
)
if direction_scope == "per layer":
direction_index = None
parameters = {}
for component in model.get_abliterable_components():
# The parameter ranges are based on experiments with various models
# and much wider ranges. They are not set in stone and might have to be
# adjusted for future models.
max_weight = trial.suggest_float(
f"{component}.max_weight",
0.8,
1.5,
)
max_weight_position = trial.suggest_float(
f"{component}.max_weight_position",
0.6 * last_layer_index,
1.0 * last_layer_index,
)
# For sampling purposes, min_weight is expressed as a fraction of max_weight,
# again because multivariate TPE doesn't support variable-range parameters.
# The value is transformed into the actual min_weight value below.
min_weight = trial.suggest_float(
f"{component}.min_weight",
0.0,
1.0,
)
min_weight_distance = trial.suggest_float(
f"{component}.min_weight_distance",
1.0,
0.6 * last_layer_index,
)
parameters[component] = AbliterationParameters(
max_weight=max_weight,
max_weight_position=max_weight_position,
min_weight=(min_weight * max_weight),
min_weight_distance=min_weight_distance,
)
trial.set_user_attr("direction_index", direction_index)
trial.set_user_attr(
"parameters", {k: asdict(v) for k, v in parameters.items()}
)
print()
print(
f"Running trial [bold]{trial_index}[/] of [bold]{settings.n_trials}[/]..."
)
print("* Parameters:")
for name, value in get_trial_parameters(trial).items():
for name, value in get_trial_parameters(settings, trial).items():
print(f" * {name} = [bold]{value}[/]")
print("* Resetting model...")
model.reset_model()
print("* Abliterating...")
model.abliterate(refusal_directions, direction_index, parameters)
if settings.use_ara_lora:
print("* Resetting model...")
model.reset_model()
print("* Abliterating (Arbitrary-Rank Ablation with LoRA)...")
model.ara_lora_abliterate(
good_module_io,
bad_module_io,
ARAParameters(**trial.user_attrs["ara_parameters"]),
)
elif settings.use_ara:
print("* Reloading model...")
model.reset_model()
print("* Abliterating (Arbitrary-Rank Ablation)...")
model.ara_abliterate(good_module_io, bad_module_io, ara_parameters)
else:
print("* Resetting model...")
model.reset_model()
print("* Abliterating...")
model.abliterate(refusal_directions, direction_index, parameters)
print("* Evaluating...")
score, kl_divergence, refusals = evaluator.get_score()
@@ -669,7 +738,7 @@ def run():
title=(
f"[Trial {trial.user_attrs['index']:>3}] "
f"Refusals: {trial.user_attrs['refusals']:>2}/{len(evaluator.bad_prompts)}, "
f"KL divergence: {trial.user_attrs['kl_divergence']:.4f}"
f"{'PIQA acc_norm' if settings.use_piqa else 'KL divergence'}: {(-1 if settings.use_piqa else 1) * trial.user_attrs['kl_divergence']:.4f}"
),
value=trial,
)
@@ -748,19 +817,38 @@ def run():
print()
print(f"Restoring model from trial [bold]{trial.user_attrs['index']}[/]...")
print("* Parameters:")
for name, value in get_trial_parameters(trial).items():
for name, value in get_trial_parameters(settings, trial).items():
print(f" * {name} = [bold]{value}[/]")
print("* Resetting model...")
model.reset_model()
print("* Abliterating...")
model.abliterate(
refusal_directions,
trial.user_attrs["direction_index"],
{
k: AbliterationParameters(**v)
for k, v in trial.user_attrs["parameters"].items()
},
)
if settings.use_ara_lora:
print("* Resetting model...")
model.reset_model()
print("* Abliterating (Arbitrary-Rank Ablation with LoRA)...")
model.ara_lora_abliterate(
good_module_io,
bad_module_io,
ARAParameters(**trial.user_attrs["ara_parameters"]),
)
elif settings.use_ara:
print("* Reloading model...")
model.reset_model()
print("* Abliterating (Arbitrary-Rank Ablation)...")
model.ara_abliterate(
good_module_io,
bad_module_io,
ARAParameters(**trial.user_attrs["ara_parameters"]),
)
else:
print("* Resetting model...")
model.reset_model()
print("* Abliterating...")
model.abliterate(
refusal_directions,
trial.user_attrs["direction_index"],
{
k: AbliterationParameters(**v)
for k, v in trial.user_attrs["parameters"].items()
},
)
while True:
print()
@@ -796,8 +884,12 @@ def run():
print("Saving LoRA adapter...")
model.model.save_pretrained(save_directory)
else:
print("Saving merged model...")
merged_model = model.get_merged_model()
if settings.use_ara:
print("Saving model...")
merged_model = model.model
else:
print("Saving merged model...")
merged_model = model.get_merged_model()
merged_model.save_pretrained(save_directory)
del merged_model
empty_cache()
@@ -851,8 +943,12 @@ def run():
token=token,
)
else:
print("Uploading merged model...")
merged_model = model.get_merged_model()
if settings.use_ara:
print("Uploading model...")
merged_model = model.model
else:
print("Uploading merged model...")
merged_model = model.get_merged_model()
merged_model.push_to_hub(
repo_id,
private=private,
@@ -891,6 +987,14 @@ def run():
card.data.tags.append("uncensored")
card.data.tags.append("decensored")
card.data.tags.append("abliterated")
if settings.use_ara:
card.data.tags.append("ara")
elif (
settings.orthogonalize_direction
and settings.row_normalization
== RowNormalization.FULL
):
card.data.tags.append("mpoa")
card.text = (
get_readme_intro(
settings,
@@ -966,6 +1070,7 @@ def run():
hflm = HFLM(
pretrained=model.model, # ty:ignore[invalid-argument-type]
tokenizer=model.tokenizer, # ty:ignore[invalid-argument-type]
batch_size="auto",
)
table = Table()
@@ -989,7 +1094,6 @@ def run():
results = lm_eval.simple_evaluate(
model=hflm,
tasks=[benchmark.task],
batch_size="auto",
)
return results["results"][benchmark.task]
+398 -16
View File
@@ -4,7 +4,7 @@
import math
from contextlib import suppress
from dataclasses import dataclass
from typing import Any, Type, cast
from typing import Any, Callable, Type, TypeAlias, cast
import bitsandbytes as bnb
import torch
@@ -14,6 +14,8 @@ from peft import LoraConfig, PeftModel, get_peft_model
from peft.tuners.lora.layer import Linear
from torch import FloatTensor, LongTensor, Tensor
from torch.nn import Module, ModuleList
from torch.optim import LBFGS
from torch.utils.hooks import RemovableHandle
from transformers import (
AutoModelForCausalLM,
AutoModelForImageTextToText,
@@ -30,7 +32,7 @@ from transformers.generation import (
)
from .config import QuantizationMethod, RowNormalization, Settings
from .utils import Prompt, batchify, empty_cache, print
from .utils import Prompt, batchify, empty_cache, mean_distances_to_knn, print
def get_model_class(
@@ -52,6 +54,23 @@ class AbliterationParameters:
min_weight_distance: float
@dataclass
class ARAParameters:
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
# 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
@@ -142,7 +161,8 @@ class Model:
if self.model is None:
raise Exception("Failed to load model with all configured dtypes.")
self._apply_lora()
if not settings.use_ara or settings.use_ara_lora:
self._apply_lora()
# LoRA B matrices are initialized to zero by default in PEFT,
# so we don't need to do anything manually.
@@ -168,21 +188,24 @@ class Model:
# because hybrid models like Qwen3.5 MoE have modules with different names
# across layers (e.g. "o_proj" on attention layers, "out_proj" on linear attention layers).
target_modules_set: set[str] = set()
module_id_to_full_name = {
id(module): module_name
for module_name, module in self.model.named_modules()
}
for layer_index, layer in enumerate(self.get_layers()):
module_id_to_leaf_name = {
id(module): module_name.split(".")[-1]
for module_name, module in layer.named_modules()
}
for layer_index in range(len(self.get_layers())):
for modules in self.get_layer_modules(layer_index).values():
for module in modules:
if id(module) in module_id_to_leaf_name:
target_modules_set.add(module_id_to_leaf_name[id(module)])
full_name = module_id_to_full_name.get(id(module))
if full_name is not None:
target_modules_set.add(full_name)
target_modules = list(target_modules_set)
target_modules = sorted(target_modules_set)
if self.settings.row_normalization != RowNormalization.FULL:
if self.settings.use_ara_lora:
lora_rank = self.settings.ara_lora_rank
elif self.settings.row_normalization != RowNormalization.FULL:
# Rank 1 is sufficient for directional ablation without renormalization.
lora_rank = 1
else:
@@ -204,7 +227,10 @@ class Model:
# so the result is a PeftModel rather than a PeftMixedModel.
self.model = cast(PeftModel, get_peft_model(self.model, self.peft_config))
print(f"* LoRA adapters initialized (targets: {', '.join(target_modules)})")
display_targets = sorted({name.rsplit(".", 1)[-1] for name in target_modules})
print(
f"* LoRA adapters initialized (target types: {', '.join(display_targets)})"
)
def _get_quantization_config(self, dtype: str) -> BitsAndBytesConfig | None:
"""
@@ -288,7 +314,11 @@ class Model:
performs full model reload with quantization config.
"""
current_model = getattr(self.model.config, "name_or_path", None)
if current_model == self.settings.model and not self.needs_reload:
if (
current_model == self.settings.model
and not self.needs_reload
and (not self.settings.use_ara or self.settings.use_ara_lora)
):
# Reset LoRA adapters to zero (identity transformation)
for name, module in self.model.named_modules():
if "lora_B" in name and hasattr(module, "weight"):
@@ -317,7 +347,8 @@ class Model:
**extra_kwargs,
)
self._apply_lora()
if not self.settings.use_ara or self.settings.use_ara_lora:
self._apply_lora()
self.needs_reload = False
@@ -341,6 +372,9 @@ class Model:
modules = {}
def try_add(component: str, module: Any):
if component not in self.settings.target_components:
return
# Only add if it's a proper nn.Module (PEFT can wrap these with LoRA)
if isinstance(module, Module):
if component not in modules:
@@ -541,6 +575,228 @@ class Model:
weight_A.data = lora_A.to(weight_A.dtype)
weight_B.data = lora_B.to(weight_B.dtype)
def ara_abliterate(
self,
good_module_io: ModuleIO,
bad_module_io: ModuleIO,
parameters: ARAParameters,
):
for layer_index in range(
parameters.start_layer_index,
parameters.end_layer_index,
):
for component, modules in self.get_layer_modules(layer_index).items():
for module_index, module in enumerate(modules):
# See above for a (partial) justification of this cast.
module = cast(Linear, module)
matrix = module.weight
row_norms = LA.vector_norm(matrix, dim=1, keepdim=True).detach()
# Helper function for reparameterization (row-norm preservation constraint).
def get_matrix() -> Tensor:
if self.settings.row_normalization == RowNormalization.FULL:
# See https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration
return row_norms * F.normalize(matrix, p=2, dim=1)
else:
return matrix
good_input, good_output = good_module_io[layer_index][component][
module_index
]
bad_input, bad_output = bad_module_io[layer_index][component][
module_index
]
good_input = good_input.to(matrix.device)
good_output = good_output.to(matrix.device)
bad_input = bad_input.to(matrix.device)
bad_output = bad_output.to(matrix.device)
def objective(matrix: Tensor) -> Tensor:
new_good_output = good_input @ matrix.T
new_bad_output = bad_input @ matrix.T
# 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
)
optimizer = LBFGS(
[matrix],
lr=1.0,
max_iter=20, # Number of internal iterations per step, *not* the number of steps.
history_size=10,
line_search_fn="strong_wolfe",
)
def closure() -> Tensor:
optimizer.zero_grad()
loss = objective(get_matrix())
loss.backward()
return loss
# Convergence usually happens within 2-3 steps, so this is more than enough.
for step in range(5):
loss = optimizer.step(closure)
# print(
# f"\\[{layer_index}/{component}/{module_index}] Step: {step}, Loss: {loss.item():.6f}"
# )
# Free the gradient buffers accumulated on the weight parameters
# during optimization. Without this, they persist on the model
# (one full-size gradient per processed weight) and can easily
# consume tens of GiB of VRAM, causing out-of-memory errors
# during the subsequent evaluation.
optimizer.zero_grad(set_to_none=True)
with torch.no_grad():
matrix.copy_(get_matrix())
def ara_lora_abliterate(
self,
good_module_io: ModuleIO,
bad_module_io: ModuleIO,
parameters: ARAParameters,
):
for layer_index in range(
parameters.start_layer_index,
parameters.end_layer_index,
):
for component, modules in self.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)
# Base weight handling and dequantization.
# 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:
# Maintain the original dequantization logic for bitsandbytes.
W_base = cast(
Tensor,
bnb.functional.dequantize_4bit(
base_weight.data,
quant_state
).to(torch.float32),
)
# Row normalization setup.
# Pre-calculate the original row norms to preserve them.
# This implements the RowNormalization.FULL logic.
W_row_norms = LA.vector_norm(W_base, dim=1, keepdim=True).detach()
# Adapter target identification.
# 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)
# Data preparation.
# Move I/O tensors to the device of the adapter weights.
good_input, good_output = good_module_io[layer_index][component][module_index]
bad_input, bad_output = 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)
# The objective function.
def objective(A: Tensor, B: Tensor) -> Tensor:
# Calculate effective weight: W_eff = W_base + B @ A.
W_eff = W_base + (B @ A)
# Apply Row Normalization (keep original norms).
if self.settings.row_normalization == RowNormalization.FULL:
# Normalize to unit length, then scale by original 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
# The original ARA loss function.
preserve_good_behavior = (
(new_good_output - good_output) ** 2
).mean()
steer_bad_behavior = (
mean_distances_to_knn(
new_bad_output,
good_output,
parameters.neighbor_count,
).mean()
+ 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
)
# Optimization loop.
# We optimize A and B, not the base matrix.
optimizer = LBFGS(
[lora_A, lora_B],
lr=1.0,
max_iter=20,
history_size=10,
line_search_fn="strong_wolfe",
)
def closure():
optimizer.zero_grad()
# Pass the actual tensors being optimized to the objective.
loss = objective(lora_A, lora_B)
loss.backward()
return loss
# Run optimization steps.
for step in range(5):
optimizer.step(closure)
# Free the gradient buffers accumulated on the LoRA adapter
# parameters during optimization (see ara_abliterate for details).
optimizer.zero_grad(set_to_none=True)
def generate(
self,
prompts: list[Prompt],
@@ -675,6 +931,132 @@ class Model:
return torch.cat(residuals, dim=0)
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
# We work with logprobs rather than probabilities for numerical stability
# when computing the KL divergence.
def get_logprobs(self, prompts: list[Prompt]) -> Tensor:
+55 -14
View File
@@ -25,8 +25,9 @@ from optuna import Trial
from psutil import Process
from questionary import Choice, Style
from rich.console import Console
from torch import Tensor
from .config import DatasetSpecification, Settings
from .config import DatasetSpecification, RowNormalization, Settings
print = Console(highlight=False).print
@@ -234,6 +235,14 @@ def batchify(items: list[T], batch_size: int) -> list[list[T]]:
return [items[i : i + batch_size] for i in range(0, len(items), batch_size)]
# 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)
def empty_cache():
# Collecting garbage is not an idempotent operation, and to avoid OOM errors,
# gc.collect() has to be called both before and after emptying the backend cache.
@@ -256,19 +265,46 @@ def empty_cache():
gc.collect()
def get_trial_parameters(trial: Trial) -> dict[str, str]:
params = {}
def get_trial_parameters(settings: Settings, trial: Trial) -> dict[str, str]:
if settings.use_ara:
parameters = trial.user_attrs["ara_parameters"]
direction_index = trial.user_attrs["direction_index"]
params["direction_index"] = (
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
)
return {
name: (f"{value:.4f}" if isinstance(value, float) else f"{value}")
for name, value in parameters.items()
}
else:
params = {}
for component, parameters in trial.user_attrs["parameters"].items():
for name, value in parameters.items():
params[f"{component}.{name}"] = f"{value:.2f}"
direction_index = trial.user_attrs["direction_index"]
params["direction_index"] = (
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
)
return params
for component, parameters in trial.user_attrs["parameters"].items():
for name, value in parameters.items():
params[f"{component}.{name}"] = f"{value:.2f}"
return params
def get_method_description(settings: Settings) -> str:
if settings.use_ara:
return (
" with the [Arbitrary-Rank Ablation (ARA)](https://github.com/p-e-w/heretic/pull/211) method"
+ (
" (with row-norm preservation)"
if settings.row_normalization == RowNormalization.FULL
else ""
)
)
elif (
settings.orthogonalize_direction
and settings.row_normalization == RowNormalization.FULL
):
return " with a variant of the [Magnitude-Preserving Orthogonal Ablation (MPOA)](https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration) method"
else:
return ""
def get_readme_intro(
@@ -285,7 +321,9 @@ def get_readme_intro(
return f"""# This is a decensored version of {
model_link
}, made using [Heretic](https://github.com/p-e-w/heretic) v{version("heretic-llm")}
}, made using [Heretic](https://github.com/p-e-w/heretic) v{version("heretic-llm")}{
get_method_description(settings)
}
## Abliteration parameters
@@ -295,7 +333,7 @@ def get_readme_intro(
chr(10).join(
[
f"| **{name}** | {value} |"
for name, value in get_trial_parameters(trial).items()
for name, value in get_trial_parameters(settings, trial).items()
]
)
}
@@ -304,7 +342,10 @@ def get_readme_intro(
| Metric | This model | Original model ({model_link}) |
| :----- | :--------: | :---------------------------: |
| **KL divergence** | {trial.user_attrs["kl_divergence"]:.4f} | 0 *(by definition)* |
| **{"PIQA acc_norm" if settings.use_piqa else "KL divergence"}** | {
(-1 if settings.use_piqa else 1) * trial.user_attrs["kl_divergence"]:.4f} | {
"*Unknown*" if settings.use_piqa else "0 *(by definition)*"
} |
| **Refusals** | {trial.user_attrs["refusals"]}/{len(bad_prompts)} | {base_refusals}/{
len(bad_prompts)
} |