diff --git a/config.default.toml b/config.default.toml index f0d6aa1..1d8588e 100644 --- a/config.default.toml +++ b/config.default.toml @@ -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.]` (and optionally `[scorer._]` 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. diff --git a/src/heretic/config.py b/src/heretic/config.py index 91a8219..8ddd4df 100644 --- a/src/heretic/config.py +++ b/src/heretic/config.py @@ -392,7 +392,7 @@ class Settings(BaseSettings): modifiers: list[ModifierConfig] = Field( default=[ ModifierConfig( - plugin="heretic.modifiers.abliteration.Abliteration", + plugin="heretic.modifiers.ara.ARA", ), ], description=( diff --git a/src/heretic/main.py b/src/heretic/main.py index 033c510..379a7f3 100644 --- a/src/heretic/main.py +++ b/src/heretic/main.py @@ -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"] ) diff --git a/src/heretic/model.py b/src/heretic/model.py index 2e63032..7112e14 100644 --- a/src/heretic/model.py +++ b/src/heretic/model.py @@ -2,12 +2,13 @@ # Copyright (C) 2025-2026 Philipp Emanuel Weidmann + 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. diff --git a/src/heretic/modifiers/abliteration.py b/src/heretic/modifiers/abliteration.py index 39b3a99..842dd42 100644 --- a/src/heretic/modifiers/abliteration.py +++ b/src/heretic/modifiers/abliteration.py @@ -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, ) diff --git a/src/heretic/modifiers/ara.py b/src/heretic/modifiers/ara.py new file mode 100644 index 0000000..f93a708 --- /dev/null +++ b/src/heretic/modifiers/ara.py @@ -0,0 +1,351 @@ +# SPDX-License-Identifier: AGPL-3.0-or-later +# Copyright (C) 2025-2026 Philipp Emanuel Weidmann + 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) diff --git a/tests/gemma-4e/config.toml b/tests/gemma-4e/config.toml index 926d293..831be6c 100644 --- a/tests/gemma-4e/config.toml +++ b/tests/gemma-4e/config.toml @@ -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 diff --git a/tests/minicpm5/config.toml b/tests/minicpm5/config.toml index d298a3a..b484c61 100644 --- a/tests/minicpm5/config.toml +++ b/tests/minicpm5/config.toml @@ -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 diff --git a/tests/mistral-3/config.toml b/tests/mistral-3/config.toml index 39bf303..50148af 100644 --- a/tests/mistral-3/config.toml +++ b/tests/mistral-3/config.toml @@ -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 diff --git a/tests/qwen2.5/config.toml b/tests/qwen2.5/config.toml index 699baba..0622f8c 100644 --- a/tests/qwen2.5/config.toml +++ b/tests/qwen2.5/config.toml @@ -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 diff --git a/tests/qwen3.5-moe/config.toml b/tests/qwen3.5-moe/config.toml index a0cd6e1..15c4993 100644 --- a/tests/qwen3.5-moe/config.toml +++ b/tests/qwen3.5-moe/config.toml @@ -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