diff --git a/src/heretic/model.py b/src/heretic/model.py index 91f5953..618bf88 100644 --- a/src/heretic/model.py +++ b/src/heretic/model.py @@ -32,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( @@ -551,9 +551,7 @@ class Model: for module_index, module in enumerate(modules): # See above for a (partial) justification of this cast. module = cast(Linear, module) - matrix = module.weight - original_matrix = matrix.detach().clone() good_input, good_output = good_module_io[layer_index][component][ module_index @@ -563,30 +561,39 @@ class Model: ] def objective(matrix: Tensor) -> Tensor: - # The results of applying the operator to inputs associated - # with "good" prompts should change as little as possible. + 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 = ( - (good_input @ matrix.T - good_output) ** 2 + (new_good_output - good_output) ** 2 ).mean() - # On average, the outputs for "bad" prompts should resemble - # the original outputs for "good" prompts (which steers the - # behavior for "bad" prompts towards that for "good" prompts). - # - # 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. + # TODO: Justify the magic weights here and make them configurable. + # Experimentally, steer_bad_behavior needs to be about an order + # of magnitude larger than preserve_good_behavior for good results. steer_bad_behavior = ( - ( - (bad_input @ matrix.T).mean(dim=0) - - good_output.mean(dim=0) - ) - ** 2 - ).mean() + 0.001 + # Pull the outputs for "bad" prompts towards + # the original outputs for "good" prompts. + * mean_distances_to_knn( + new_bad_output, + 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 ( preserve_good_behavior_weight * preserve_good_behavior diff --git a/src/heretic/utils.py b/src/heretic/utils.py index dd3a176..c9c8951 100644 --- a/src/heretic/utils.py +++ b/src/heretic/utils.py @@ -25,6 +25,7 @@ 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 @@ -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)] +# 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.