mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-28 23:11:25 -07:00
feat(ara): implement matrix optimization
This commit is contained in:
+20
-2
@@ -207,7 +207,7 @@ def run():
|
|||||||
)
|
)
|
||||||
|
|
||||||
# We don't need gradients as we only do inference.
|
# We don't need gradients as we only do inference.
|
||||||
torch.set_grad_enabled(False)
|
# torch.set_grad_enabled(False)
|
||||||
|
|
||||||
# While determining the optimal batch size, we will try many different batch sizes,
|
# 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)
|
# resulting in many computation graphs being compiled. Raising the limit (default = 8)
|
||||||
@@ -422,7 +422,25 @@ def run():
|
|||||||
print("Obtaining module I/O for bad prompts...")
|
print("Obtaining module I/O for bad prompts...")
|
||||||
bad_module_io = model.get_module_io_batched(bad_prompts)
|
bad_module_io = model.get_module_io_batched(bad_prompts)
|
||||||
|
|
||||||
print(good_module_io)
|
# print(good_module_io)
|
||||||
|
|
||||||
|
print()
|
||||||
|
print("Performing Arbitrary-Rank Ablation...")
|
||||||
|
|
||||||
|
model.ara_abliterate(
|
||||||
|
good_module_io,
|
||||||
|
bad_module_io,
|
||||||
|
0,
|
||||||
|
len(model.get_layers()),
|
||||||
|
1.0,
|
||||||
|
1.0,
|
||||||
|
1.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
print()
|
||||||
|
print("Evaluating...")
|
||||||
|
evaluator.get_score()
|
||||||
|
return
|
||||||
|
|
||||||
print()
|
print()
|
||||||
print("Calculating per-layer refusal directions...")
|
print("Calculating per-layer refusal directions...")
|
||||||
|
|||||||
+93
-7
@@ -4,7 +4,7 @@
|
|||||||
import math
|
import math
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Callable, Type, cast
|
from typing import Any, Callable, Type, TypeAlias, cast
|
||||||
|
|
||||||
import bitsandbytes as bnb
|
import bitsandbytes as bnb
|
||||||
import torch
|
import torch
|
||||||
@@ -14,6 +14,7 @@ from peft import LoraConfig, PeftModel, get_peft_model
|
|||||||
from peft.tuners.lora.layer import Linear
|
from peft.tuners.lora.layer import Linear
|
||||||
from torch import FloatTensor, LongTensor, Tensor
|
from torch import FloatTensor, LongTensor, Tensor
|
||||||
from torch.nn import Module, ModuleList
|
from torch.nn import Module, ModuleList
|
||||||
|
from torch.optim import LBFGS
|
||||||
from torch.utils.hooks import RemovableHandle
|
from torch.utils.hooks import RemovableHandle
|
||||||
from transformers import (
|
from transformers import (
|
||||||
AutoModelForCausalLM,
|
AutoModelForCausalLM,
|
||||||
@@ -53,6 +54,13 @@ class AbliterationParameters:
|
|||||||
min_weight_distance: float
|
min_weight_distance: float
|
||||||
|
|
||||||
|
|
||||||
|
# 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:
|
class Model:
|
||||||
model: PreTrainedModel | PeftModel
|
model: PreTrainedModel | PeftModel
|
||||||
tokenizer: PreTrainedTokenizerBase
|
tokenizer: PreTrainedTokenizerBase
|
||||||
@@ -145,7 +153,7 @@ class Model:
|
|||||||
if self.model is None:
|
if self.model is None:
|
||||||
raise Exception("Failed to load model with all configured dtypes.")
|
raise Exception("Failed to load model with all configured dtypes.")
|
||||||
|
|
||||||
self._apply_lora()
|
# self._apply_lora()
|
||||||
|
|
||||||
# LoRA B matrices are initialized to zero by default in PEFT,
|
# LoRA B matrices are initialized to zero by default in PEFT,
|
||||||
# so we don't need to do anything manually.
|
# so we don't need to do anything manually.
|
||||||
@@ -520,6 +528,84 @@ class Model:
|
|||||||
weight_A.data = lora_A.to(weight_A.dtype)
|
weight_A.data = lora_A.to(weight_A.dtype)
|
||||||
weight_B.data = lora_B.to(weight_B.dtype)
|
weight_B.data = lora_B.to(weight_B.dtype)
|
||||||
|
|
||||||
|
def ara_abliterate(
|
||||||
|
self,
|
||||||
|
good_module_io: ModuleIO,
|
||||||
|
bad_module_io: ModuleIO,
|
||||||
|
start_layer_index: int,
|
||||||
|
end_layer_index: int,
|
||||||
|
preserve_good_behavior_weight: float,
|
||||||
|
steer_bad_behavior_weight: float,
|
||||||
|
tie_to_original_matrix_weight: float,
|
||||||
|
):
|
||||||
|
for layer_index in range(start_layer_index, 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
|
||||||
|
original_matrix = matrix.detach().clone()
|
||||||
|
|
||||||
|
good_input, good_output = good_module_io[layer_index][component][
|
||||||
|
module_index
|
||||||
|
]
|
||||||
|
bad_input, bad_output = bad_module_io[layer_index][component][
|
||||||
|
module_index
|
||||||
|
]
|
||||||
|
|
||||||
|
def objective(matrix: Tensor) -> Tensor:
|
||||||
|
# The results of applying the operator to inputs associated
|
||||||
|
# with "good" prompts should change as little as possible.
|
||||||
|
preserve_good_behavior = (
|
||||||
|
(good_input @ matrix.T - 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).
|
||||||
|
steer_bad_behavior = (
|
||||||
|
(
|
||||||
|
(bad_input @ matrix.T).mean(dim=0)
|
||||||
|
- good_output.mean(dim=0)
|
||||||
|
)
|
||||||
|
** 2
|
||||||
|
).mean()
|
||||||
|
|
||||||
|
# The matrix itself should change as little as possible overall.
|
||||||
|
# This prevents overfitting due to underdetermination of the
|
||||||
|
# optimization problem from a relatively small number of I/O pairs.
|
||||||
|
tie_to_original_matrix = (
|
||||||
|
(matrix - original_matrix) ** 2
|
||||||
|
).mean()
|
||||||
|
|
||||||
|
return (
|
||||||
|
preserve_good_behavior_weight * preserve_good_behavior
|
||||||
|
+ steer_bad_behavior_weight * steer_bad_behavior
|
||||||
|
+ tie_to_original_matrix_weight * tie_to_original_matrix
|
||||||
|
)
|
||||||
|
|
||||||
|
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(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}"
|
||||||
|
)
|
||||||
|
|
||||||
def generate(
|
def generate(
|
||||||
self,
|
self,
|
||||||
prompts: list[Prompt],
|
prompts: list[Prompt],
|
||||||
@@ -657,12 +743,12 @@ class Model:
|
|||||||
def get_module_io(
|
def get_module_io(
|
||||||
self,
|
self,
|
||||||
prompts: list[Prompt],
|
prompts: list[Prompt],
|
||||||
) -> list[dict[str, dict[int, tuple[Tensor, Tensor]]]]:
|
) -> ModuleIO:
|
||||||
# The list contains one element per layer.
|
# The list contains one element per layer.
|
||||||
# Each element maps from the component name to a (possibly sparse) mapping
|
# 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
|
# from the module index to an (input, output) tuple containing the I/O
|
||||||
# tensors of shape (prompt, component).
|
# tensors of shape (prompt, component).
|
||||||
module_io: list[dict[str, dict[int, tuple[Tensor, Tensor]]]] = []
|
module_io: ModuleIO = []
|
||||||
|
|
||||||
def get_hook(
|
def get_hook(
|
||||||
layer_index: int,
|
layer_index: int,
|
||||||
@@ -722,19 +808,19 @@ class Model:
|
|||||||
def get_module_io_batched(
|
def get_module_io_batched(
|
||||||
self,
|
self,
|
||||||
prompts: list[Prompt],
|
prompts: list[Prompt],
|
||||||
) -> list[dict[str, dict[int, tuple[Tensor, Tensor]]]]:
|
) -> ModuleIO:
|
||||||
# Aggregating batch results is more complicated for module I/O
|
# Aggregating batch results is more complicated for module I/O
|
||||||
# than for other get_*_batched methods, because the structure of the results
|
# than for other get_*_batched methods, because the structure of the results
|
||||||
# might differ between batches, as whether individual modules activate
|
# might differ between batches, as whether individual modules activate
|
||||||
# can depend on the prompt (in particular for MoE models).
|
# can depend on the prompt (in particular for MoE models).
|
||||||
# In practice, inhomogeneous results should be very rare, but to be fully
|
# In practice, inhomogeneous results should be very rare, but to be fully
|
||||||
# generic, this logic is required.
|
# generic, this logic is required.
|
||||||
module_io_batches = [
|
module_io_batches: list[ModuleIO] = [
|
||||||
self.get_module_io(batch)
|
self.get_module_io(batch)
|
||||||
for batch in batchify(prompts, self.settings.batch_size)
|
for batch in batchify(prompts, self.settings.batch_size)
|
||||||
]
|
]
|
||||||
|
|
||||||
module_io: list[dict[str, dict[int, tuple[Tensor, Tensor]]]] = []
|
module_io: ModuleIO = []
|
||||||
|
|
||||||
for layer_index in range(len(self.get_layers())):
|
for layer_index in range(len(self.get_layers())):
|
||||||
module_io.append({})
|
module_io.append({})
|
||||||
|
|||||||
Reference in New Issue
Block a user