mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-28 23:11:25 -07:00
feat(ara): improve steering term
This commit is contained in:
+30
-23
@@ -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
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user