feat(ara): improve steering term

This commit is contained in:
Philipp Emanuel Weidmann
2026-03-04 11:06:01 +05:30
parent bd1fa0ade4
commit 56e57adf36
2 changed files with 39 additions and 23 deletions
+30 -23
View File
@@ -32,7 +32,7 @@ from transformers.generation import (
) )
from .config import QuantizationMethod, RowNormalization, Settings 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( def get_model_class(
@@ -551,9 +551,7 @@ class Model:
for module_index, module in enumerate(modules): for module_index, module in enumerate(modules):
# See above for a (partial) justification of this cast. # See above for a (partial) justification of this cast.
module = cast(Linear, module) module = cast(Linear, module)
matrix = module.weight matrix = module.weight
original_matrix = matrix.detach().clone()
good_input, good_output = good_module_io[layer_index][component][ good_input, good_output = good_module_io[layer_index][component][
module_index module_index
@@ -563,30 +561,39 @@ class Model:
] ]
def objective(matrix: Tensor) -> Tensor: def objective(matrix: Tensor) -> Tensor:
# The results of applying the operator to inputs associated new_good_output = good_input @ matrix.T
# with "good" prompts should change as little as possible. new_bad_output = bad_input @ matrix.T
# The outputs for "good" prompts should change as little as possible.
preserve_good_behavior = ( preserve_good_behavior = (
(good_input @ matrix.T - good_output) ** 2 (new_good_output - good_output) ** 2
).mean() ).mean()
# On average, the outputs for "bad" prompts should resemble # TODO: Justify the magic weights here and make them configurable.
# the original outputs for "good" prompts (which steers the # Experimentally, steer_bad_behavior needs to be about an order
# behavior for "bad" prompts towards that for "good" prompts). # of magnitude larger than preserve_good_behavior for good results.
#
# TODO: An alternative formulation could use the mean distance
# of "bad" outputs from the boundary of the core cluster
# of original "good" outputs. This would classify an output
# configuration as optimal as long as all "bad" outputs
# are inside the same cluster as the "good" outputs, even
# if their centroid is different from those of the "good"
# outputs.
steer_bad_behavior = ( steer_bad_behavior = (
( 0.001
(bad_input @ matrix.T).mean(dim=0) # Pull the outputs for "bad" prompts towards
- good_output.mean(dim=0) # the original outputs for "good" prompts.
) * mean_distances_to_knn(
** 2 new_bad_output,
).mean() good_output,
10,
).mean()
+ 0.0001
# 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.
* -mean_distances_to_knn(
new_bad_output,
bad_output,
10,
).mean()
)
return ( return (
preserve_good_behavior_weight * preserve_good_behavior preserve_good_behavior_weight * preserve_good_behavior
+9
View File
@@ -25,6 +25,7 @@ from optuna import Trial
from psutil import Process from psutil import Process
from questionary import Choice, Style from questionary import Choice, Style
from rich.console import Console from rich.console import Console
from torch import Tensor
from .config import DatasetSpecification, Settings from .config import DatasetSpecification, Settings
@@ -228,6 +229,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)] 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(): def empty_cache():
# Collecting garbage is not an idempotent operation, and to avoid OOM errors, # 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. # gc.collect() has to be called both before and after emptying the backend cache.