mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-10 06:09:08 -07:00
Compare commits
19 Commits
modifier-plugins
...
ara
| Author | SHA1 | Date | |
|---|---|---|---|
| edc3b12345 | |||
| 25979ad7d0 | |||
| 3b70fe5dfa | |||
| f7a456bd0c | |||
| 988c6bd90e | |||
| c925f5e802 | |||
| 4a6304c361 | |||
| c76416fe03 | |||
| 2bb203ee47 | |||
| d79a443e6f | |||
| 0bb9521fbe | |||
| 992fb3a4b3 | |||
| 304c14adc7 | |||
| 56e57adf36 | |||
| bd1fa0ade4 | |||
| 3c5d6920bf | |||
| b8f4a9c985 | |||
| 154241f8a2 | |||
| ea7c59a55a |
+37
-1
@@ -188,6 +188,42 @@ class Settings(BaseSettings):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
target_components: list[str] = Field(
|
||||||
|
default=["attn.o_proj", "mlp.down_proj"],
|
||||||
|
description=(
|
||||||
|
"List of component names to target for abliteration. "
|
||||||
|
'Currently supported values are "attn.o_proj" and "mlp.down_proj".'
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
use_ara: bool = Field(
|
||||||
|
default=True,
|
||||||
|
description=(
|
||||||
|
"Whether to use Arbitrary-Rank Ablation (ARA), an abliteration method based on matrix optimization, "
|
||||||
|
"instead of traditional directional ablation."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
use_ara_lora: bool = Field(
|
||||||
|
default=False,
|
||||||
|
description=(
|
||||||
|
"Use LoRA in ARA instead of full-weight editing. Makes it compatible with quantization and removes model reloads."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
ara_lora_rank: int = Field(
|
||||||
|
default=128,
|
||||||
|
description="If LoRA is used in ARA, this sets up its rank. Keep it high enough to simulate the 'arbitrary' effect.",
|
||||||
|
)
|
||||||
|
|
||||||
|
use_piqa: bool = Field(
|
||||||
|
default=False,
|
||||||
|
description=(
|
||||||
|
"Whether to use the Physical Interaction: Question Answering (PIQA) benchmark "
|
||||||
|
"as the quality metric instead of the Kullback-Leibler divergence."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
orthogonalize_direction: bool = Field(
|
orthogonalize_direction: bool = Field(
|
||||||
default=False,
|
default=False,
|
||||||
description=(
|
description=(
|
||||||
@@ -197,7 +233,7 @@ class Settings(BaseSettings):
|
|||||||
)
|
)
|
||||||
|
|
||||||
row_normalization: RowNormalization = Field(
|
row_normalization: RowNormalization = Field(
|
||||||
default=RowNormalization.NONE,
|
default=RowNormalization.FULL,
|
||||||
description=(
|
description=(
|
||||||
"How to apply row normalization of the weights. Options: "
|
"How to apply row normalization of the weights. Options: "
|
||||||
'"none" (no normalization), '
|
'"none" (no normalization), '
|
||||||
|
|||||||
+53
-28
@@ -1,7 +1,9 @@
|
|||||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||||
|
|
||||||
|
import lm_eval
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
from lm_eval.models.huggingface import HFLM
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from .config import Settings
|
from .config import Settings
|
||||||
@@ -21,15 +23,16 @@ class Evaluator:
|
|||||||
self.settings = settings
|
self.settings = settings
|
||||||
self.model = model
|
self.model = model
|
||||||
|
|
||||||
print()
|
if not settings.use_piqa:
|
||||||
print(
|
print()
|
||||||
f"Loading good evaluation prompts from [bold]{settings.good_evaluation_prompts.dataset}[/]..."
|
print(
|
||||||
)
|
f"Loading good evaluation prompts from [bold]{settings.good_evaluation_prompts.dataset}[/]..."
|
||||||
self.good_prompts = load_prompts(settings, settings.good_evaluation_prompts)
|
)
|
||||||
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
|
self.good_prompts = load_prompts(settings, settings.good_evaluation_prompts)
|
||||||
|
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
|
||||||
|
|
||||||
print("* Obtaining first-token probability distributions...")
|
print("* Obtaining first-token probability distributions...")
|
||||||
self.base_logprobs = model.get_logprobs_batched(self.good_prompts)
|
self.base_logprobs = model.get_logprobs_batched(self.good_prompts)
|
||||||
|
|
||||||
print()
|
print()
|
||||||
print(
|
print(
|
||||||
@@ -93,35 +96,57 @@ class Evaluator:
|
|||||||
return refusal_count
|
return refusal_count
|
||||||
|
|
||||||
def get_score(self) -> tuple[tuple[float, float], float, int]:
|
def get_score(self) -> tuple[tuple[float, float], float, int]:
|
||||||
print(" * Obtaining first-token probability distributions...")
|
if self.settings.use_piqa:
|
||||||
logprobs = self.model.get_logprobs_batched(self.good_prompts)
|
print(" * Running PIQA benchmark...")
|
||||||
kl_divergence = F.kl_div(
|
hflm = HFLM(
|
||||||
logprobs,
|
pretrained=self.model.model, # ty:ignore[invalid-argument-type]
|
||||||
self.base_logprobs,
|
tokenizer=self.model.tokenizer, # ty:ignore[invalid-argument-type]
|
||||||
reduction="batchmean",
|
batch_size="auto",
|
||||||
log_target=True,
|
)
|
||||||
).item()
|
results = lm_eval.simple_evaluate(
|
||||||
print(f" * KL divergence: [bold]{kl_divergence:.4f}[/]")
|
model=hflm,
|
||||||
|
tasks=["piqa"],
|
||||||
|
)
|
||||||
|
piqa_acc_norm: float = results["results"]["piqa"]["acc_norm,none"]
|
||||||
|
print(f" * PIQA acc_norm: [bold]{piqa_acc_norm:.4f}[/]")
|
||||||
|
else:
|
||||||
|
print(" * Obtaining first-token probability distributions...")
|
||||||
|
logprobs = self.model.get_logprobs_batched(self.good_prompts)
|
||||||
|
kl_divergence = F.kl_div(
|
||||||
|
logprobs,
|
||||||
|
self.base_logprobs,
|
||||||
|
reduction="batchmean",
|
||||||
|
log_target=True,
|
||||||
|
).item()
|
||||||
|
print(f" * KL divergence: [bold]{kl_divergence:.4f}[/]")
|
||||||
|
|
||||||
print(" * Counting model refusals...")
|
print(" * Counting model refusals...")
|
||||||
refusals = self.count_refusals()
|
refusals = self.count_refusals()
|
||||||
print(f" * Refusals: [bold]{refusals}[/]/{len(self.bad_prompts)}")
|
print(f" * Refusals: [bold]{refusals}[/]/{len(self.bad_prompts)}")
|
||||||
|
|
||||||
kl_divergence_scale = self.settings.kl_divergence_scale
|
|
||||||
kl_divergence_target = self.settings.kl_divergence_target
|
|
||||||
|
|
||||||
refusals_score = (
|
refusals_score = (
|
||||||
refusals / self.base_refusals if self.base_refusals > 0 else float(refusals)
|
refusals / self.base_refusals if self.base_refusals > 0 else float(refusals)
|
||||||
)
|
)
|
||||||
|
|
||||||
if kl_divergence >= kl_divergence_target:
|
if self.settings.use_piqa:
|
||||||
kld_score = kl_divergence / kl_divergence_scale
|
score = (
|
||||||
|
-piqa_acc_norm,
|
||||||
|
refusals_score,
|
||||||
|
)
|
||||||
|
|
||||||
|
return score, -piqa_acc_norm, refusals
|
||||||
else:
|
else:
|
||||||
kld_score = refusals_score * kl_divergence_target / kl_divergence_scale
|
kl_divergence_scale = self.settings.kl_divergence_scale
|
||||||
|
kl_divergence_target = self.settings.kl_divergence_target
|
||||||
|
|
||||||
score = (
|
if kl_divergence >= kl_divergence_target:
|
||||||
kld_score,
|
kld_score = kl_divergence / kl_divergence_scale
|
||||||
refusals_score,
|
else:
|
||||||
)
|
kld_score = refusals_score * kl_divergence_target / kl_divergence_scale
|
||||||
|
|
||||||
return score, kl_divergence, refusals
|
score = (
|
||||||
|
kld_score,
|
||||||
|
refusals_score,
|
||||||
|
)
|
||||||
|
|
||||||
|
return score, kl_divergence, refusals
|
||||||
|
|||||||
+213
-109
@@ -51,9 +51,9 @@ from rich.table import Table
|
|||||||
from rich.traceback import install
|
from rich.traceback import install
|
||||||
|
|
||||||
from .analyzer import Analyzer
|
from .analyzer import Analyzer
|
||||||
from .config import QuantizationMethod, Settings
|
from .config import QuantizationMethod, RowNormalization, Settings
|
||||||
from .evaluator import Evaluator
|
from .evaluator import Evaluator
|
||||||
from .model import AbliterationParameters, Model, get_model_class
|
from .model import AbliterationParameters, ARAParameters, Model, get_model_class
|
||||||
from .utils import (
|
from .utils import (
|
||||||
empty_cache,
|
empty_cache,
|
||||||
format_duration,
|
format_duration,
|
||||||
@@ -227,8 +227,9 @@ def run():
|
|||||||
"[bold yellow]No GPU or other accelerator detected. Operations will be slow.[/]"
|
"[bold yellow]No GPU or other accelerator detected. Operations will be slow.[/]"
|
||||||
)
|
)
|
||||||
|
|
||||||
# We don't need gradients as we only do inference.
|
if not settings.use_ara:
|
||||||
torch.set_grad_enabled(False)
|
# We don't need gradients as we only do inference.
|
||||||
|
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)
|
||||||
@@ -451,40 +452,47 @@ def run():
|
|||||||
evaluator.get_score()
|
evaluator.get_score()
|
||||||
return
|
return
|
||||||
|
|
||||||
print()
|
if settings.use_ara:
|
||||||
print("Calculating per-layer refusal directions...")
|
print()
|
||||||
print("* Obtaining residuals for good prompts...")
|
print("Obtaining module I/O for good prompts...")
|
||||||
good_residuals = model.get_residuals_batched(good_prompts)
|
good_module_io = model.get_module_io_batched(good_prompts)
|
||||||
print("* Obtaining residuals for bad prompts...")
|
print("Obtaining module I/O for bad prompts...")
|
||||||
bad_residuals = model.get_residuals_batched(bad_prompts)
|
bad_module_io = model.get_module_io_batched(bad_prompts)
|
||||||
|
else:
|
||||||
|
print()
|
||||||
|
print("Calculating per-layer refusal directions...")
|
||||||
|
print("* Obtaining residuals for good prompts...")
|
||||||
|
good_residuals = model.get_residuals_batched(good_prompts)
|
||||||
|
print("* Obtaining residuals for bad prompts...")
|
||||||
|
bad_residuals = model.get_residuals_batched(bad_prompts)
|
||||||
|
|
||||||
good_means = good_residuals.mean(dim=0)
|
good_means = good_residuals.mean(dim=0)
|
||||||
bad_means = bad_residuals.mean(dim=0)
|
bad_means = bad_residuals.mean(dim=0)
|
||||||
|
|
||||||
refusal_directions = F.normalize(bad_means - good_means, p=2, dim=1)
|
refusal_directions = F.normalize(bad_means - good_means, p=2, dim=1)
|
||||||
|
|
||||||
if settings.orthogonalize_direction:
|
if settings.orthogonalize_direction:
|
||||||
# Implements https://huggingface.co/blog/grimjim/projected-abliteration
|
# Implements https://huggingface.co/blog/grimjim/projected-abliteration
|
||||||
# Adjust the refusal directions so that only the component that is
|
# Adjust the refusal directions so that only the component that is
|
||||||
# orthogonal to the good direction is subtracted during abliteration.
|
# orthogonal to the good direction is subtracted during abliteration.
|
||||||
good_directions = F.normalize(good_means, p=2, dim=1)
|
good_directions = F.normalize(good_means, p=2, dim=1)
|
||||||
projection_vector = torch.sum(refusal_directions * good_directions, dim=1)
|
projection_vector = torch.sum(refusal_directions * good_directions, dim=1)
|
||||||
refusal_directions = (
|
refusal_directions = (
|
||||||
refusal_directions - projection_vector.unsqueeze(1) * good_directions
|
refusal_directions - projection_vector.unsqueeze(1) * good_directions
|
||||||
)
|
)
|
||||||
refusal_directions = F.normalize(refusal_directions, p=2, dim=1)
|
refusal_directions = F.normalize(refusal_directions, p=2, dim=1)
|
||||||
|
|
||||||
analyzer = Analyzer(settings, model, good_residuals, bad_residuals)
|
analyzer = Analyzer(settings, model, good_residuals, bad_residuals)
|
||||||
|
|
||||||
if settings.print_residual_geometry:
|
if settings.print_residual_geometry:
|
||||||
analyzer.print_residual_geometry()
|
analyzer.print_residual_geometry()
|
||||||
|
|
||||||
if settings.plot_residuals:
|
if settings.plot_residuals:
|
||||||
analyzer.plot_residuals()
|
analyzer.plot_residuals()
|
||||||
|
|
||||||
# We don't need the residuals after computing refusal directions.
|
# We don't need the residuals after computing refusal directions.
|
||||||
del good_residuals, bad_residuals, analyzer
|
del good_residuals, bad_residuals, analyzer
|
||||||
empty_cache()
|
empty_cache()
|
||||||
|
|
||||||
trial_index = 0
|
trial_index = 0
|
||||||
start_index = 0
|
start_index = 0
|
||||||
@@ -495,83 +503,144 @@ def run():
|
|||||||
trial_index += 1
|
trial_index += 1
|
||||||
trial.set_user_attr("index", trial_index)
|
trial.set_user_attr("index", trial_index)
|
||||||
|
|
||||||
direction_scope = trial.suggest_categorical(
|
if settings.use_ara:
|
||||||
"direction_scope",
|
start_layer_index = trial.suggest_int(
|
||||||
[
|
"start_layer_index",
|
||||||
"global",
|
0,
|
||||||
"per layer",
|
len(model.get_layers()) // 2,
|
||||||
],
|
|
||||||
)
|
|
||||||
|
|
||||||
last_layer_index = len(model.get_layers()) - 1
|
|
||||||
|
|
||||||
# Discrimination between "harmful" and "harmless" inputs is usually strongest
|
|
||||||
# in layers slightly past the midpoint of the layer stack. See the original
|
|
||||||
# abliteration paper (https://arxiv.org/abs/2406.11717) for a deeper analysis.
|
|
||||||
#
|
|
||||||
# Note that we always sample this parameter even though we only need it for
|
|
||||||
# the "global" direction scope. The reason is that multivariate TPE doesn't
|
|
||||||
# work with conditional or variable-range parameters.
|
|
||||||
direction_index = trial.suggest_float(
|
|
||||||
"direction_index",
|
|
||||||
0.4 * last_layer_index,
|
|
||||||
0.9 * last_layer_index,
|
|
||||||
)
|
|
||||||
|
|
||||||
if direction_scope == "per layer":
|
|
||||||
direction_index = None
|
|
||||||
|
|
||||||
parameters = {}
|
|
||||||
|
|
||||||
for component in model.get_abliterable_components():
|
|
||||||
# The parameter ranges are based on experiments with various models
|
|
||||||
# and much wider ranges. They are not set in stone and might have to be
|
|
||||||
# adjusted for future models.
|
|
||||||
max_weight = trial.suggest_float(
|
|
||||||
f"{component}.max_weight",
|
|
||||||
0.8,
|
|
||||||
1.5,
|
|
||||||
)
|
)
|
||||||
max_weight_position = trial.suggest_float(
|
end_layer_index = trial.suggest_int(
|
||||||
f"{component}.max_weight_position",
|
"end_layer_index",
|
||||||
0.6 * last_layer_index,
|
len(model.get_layers()) // 2,
|
||||||
1.0 * last_layer_index,
|
len(model.get_layers()),
|
||||||
)
|
)
|
||||||
# For sampling purposes, min_weight is expressed as a fraction of max_weight,
|
preserve_good_behavior_weight = trial.suggest_float(
|
||||||
# again because multivariate TPE doesn't support variable-range parameters.
|
"preserve_good_behavior_weight",
|
||||||
# The value is transformed into the actual min_weight value below.
|
|
||||||
min_weight = trial.suggest_float(
|
|
||||||
f"{component}.min_weight",
|
|
||||||
0.0,
|
0.0,
|
||||||
1.0,
|
1.0,
|
||||||
)
|
)
|
||||||
min_weight_distance = trial.suggest_float(
|
steer_bad_behavior_weight = trial.suggest_float(
|
||||||
f"{component}.min_weight_distance",
|
"steer_bad_behavior_weight",
|
||||||
|
0.0001,
|
||||||
1.0,
|
1.0,
|
||||||
0.6 * last_layer_index,
|
log=True,
|
||||||
|
)
|
||||||
|
overcorrect_relative_weight = trial.suggest_float(
|
||||||
|
"overcorrect_relative_weight",
|
||||||
|
0.0,
|
||||||
|
1.3,
|
||||||
|
)
|
||||||
|
neighbor_count = trial.suggest_int(
|
||||||
|
"neighbor_count",
|
||||||
|
1,
|
||||||
|
15,
|
||||||
)
|
)
|
||||||
|
|
||||||
parameters[component] = AbliterationParameters(
|
ara_parameters = ARAParameters(
|
||||||
max_weight=max_weight,
|
start_layer_index=start_layer_index,
|
||||||
max_weight_position=max_weight_position,
|
end_layer_index=end_layer_index,
|
||||||
min_weight=(min_weight * max_weight),
|
preserve_good_behavior_weight=preserve_good_behavior_weight,
|
||||||
min_weight_distance=min_weight_distance,
|
steer_bad_behavior_weight=steer_bad_behavior_weight,
|
||||||
|
overcorrect_relative_weight=overcorrect_relative_weight,
|
||||||
|
neighbor_count=neighbor_count,
|
||||||
)
|
)
|
||||||
|
|
||||||
trial.set_user_attr("direction_index", direction_index)
|
trial.set_user_attr("ara_parameters", asdict(ara_parameters))
|
||||||
trial.set_user_attr("parameters", {k: asdict(v) for k, v in parameters.items()})
|
else:
|
||||||
|
direction_scope = trial.suggest_categorical(
|
||||||
|
"direction_scope",
|
||||||
|
[
|
||||||
|
"global",
|
||||||
|
"per layer",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
last_layer_index = len(model.get_layers()) - 1
|
||||||
|
|
||||||
|
# Discrimination between "harmful" and "harmless" inputs is usually strongest
|
||||||
|
# in layers slightly past the midpoint of the layer stack. See the original
|
||||||
|
# abliteration paper (https://arxiv.org/abs/2406.11717) for a deeper analysis.
|
||||||
|
#
|
||||||
|
# Note that we always sample this parameter even though we only need it for
|
||||||
|
# the "global" direction scope. The reason is that multivariate TPE doesn't
|
||||||
|
# work with conditional or variable-range parameters.
|
||||||
|
direction_index = trial.suggest_float(
|
||||||
|
"direction_index",
|
||||||
|
0.4 * last_layer_index,
|
||||||
|
0.9 * last_layer_index,
|
||||||
|
)
|
||||||
|
|
||||||
|
if direction_scope == "per layer":
|
||||||
|
direction_index = None
|
||||||
|
|
||||||
|
parameters = {}
|
||||||
|
|
||||||
|
for component in model.get_abliterable_components():
|
||||||
|
# The parameter ranges are based on experiments with various models
|
||||||
|
# and much wider ranges. They are not set in stone and might have to be
|
||||||
|
# adjusted for future models.
|
||||||
|
max_weight = trial.suggest_float(
|
||||||
|
f"{component}.max_weight",
|
||||||
|
0.8,
|
||||||
|
1.5,
|
||||||
|
)
|
||||||
|
max_weight_position = trial.suggest_float(
|
||||||
|
f"{component}.max_weight_position",
|
||||||
|
0.6 * last_layer_index,
|
||||||
|
1.0 * last_layer_index,
|
||||||
|
)
|
||||||
|
# For sampling purposes, min_weight is expressed as a fraction of max_weight,
|
||||||
|
# again because multivariate TPE doesn't support variable-range parameters.
|
||||||
|
# The value is transformed into the actual min_weight value below.
|
||||||
|
min_weight = trial.suggest_float(
|
||||||
|
f"{component}.min_weight",
|
||||||
|
0.0,
|
||||||
|
1.0,
|
||||||
|
)
|
||||||
|
min_weight_distance = trial.suggest_float(
|
||||||
|
f"{component}.min_weight_distance",
|
||||||
|
1.0,
|
||||||
|
0.6 * last_layer_index,
|
||||||
|
)
|
||||||
|
|
||||||
|
parameters[component] = AbliterationParameters(
|
||||||
|
max_weight=max_weight,
|
||||||
|
max_weight_position=max_weight_position,
|
||||||
|
min_weight=(min_weight * max_weight),
|
||||||
|
min_weight_distance=min_weight_distance,
|
||||||
|
)
|
||||||
|
|
||||||
|
trial.set_user_attr("direction_index", direction_index)
|
||||||
|
trial.set_user_attr(
|
||||||
|
"parameters", {k: asdict(v) for k, v in parameters.items()}
|
||||||
|
)
|
||||||
|
|
||||||
print()
|
print()
|
||||||
print(
|
print(
|
||||||
f"Running trial [bold]{trial_index}[/] of [bold]{settings.n_trials}[/]..."
|
f"Running trial [bold]{trial_index}[/] of [bold]{settings.n_trials}[/]..."
|
||||||
)
|
)
|
||||||
print("* Parameters:")
|
print("* Parameters:")
|
||||||
for name, value in get_trial_parameters(trial).items():
|
for name, value in get_trial_parameters(settings, trial).items():
|
||||||
print(f" * {name} = [bold]{value}[/]")
|
print(f" * {name} = [bold]{value}[/]")
|
||||||
print("* Resetting model...")
|
if settings.use_ara_lora:
|
||||||
model.reset_model()
|
print("* Resetting model...")
|
||||||
print("* Abliterating...")
|
model.reset_model()
|
||||||
model.abliterate(refusal_directions, direction_index, parameters)
|
print("* Abliterating (Arbitrary-Rank Ablation with LoRA)...")
|
||||||
|
model.ara_lora_abliterate(
|
||||||
|
good_module_io,
|
||||||
|
bad_module_io,
|
||||||
|
ARAParameters(**trial.user_attrs["ara_parameters"]),
|
||||||
|
)
|
||||||
|
elif settings.use_ara:
|
||||||
|
print("* Reloading model...")
|
||||||
|
model.reset_model()
|
||||||
|
print("* Abliterating (Arbitrary-Rank Ablation)...")
|
||||||
|
model.ara_abliterate(good_module_io, bad_module_io, ara_parameters)
|
||||||
|
else:
|
||||||
|
print("* Resetting model...")
|
||||||
|
model.reset_model()
|
||||||
|
print("* Abliterating...")
|
||||||
|
model.abliterate(refusal_directions, direction_index, parameters)
|
||||||
print("* Evaluating...")
|
print("* Evaluating...")
|
||||||
score, kl_divergence, refusals = evaluator.get_score()
|
score, kl_divergence, refusals = evaluator.get_score()
|
||||||
|
|
||||||
@@ -669,7 +738,7 @@ def run():
|
|||||||
title=(
|
title=(
|
||||||
f"[Trial {trial.user_attrs['index']:>3}] "
|
f"[Trial {trial.user_attrs['index']:>3}] "
|
||||||
f"Refusals: {trial.user_attrs['refusals']:>2}/{len(evaluator.bad_prompts)}, "
|
f"Refusals: {trial.user_attrs['refusals']:>2}/{len(evaluator.bad_prompts)}, "
|
||||||
f"KL divergence: {trial.user_attrs['kl_divergence']:.4f}"
|
f"{'PIQA acc_norm' if settings.use_piqa else 'KL divergence'}: {(-1 if settings.use_piqa else 1) * trial.user_attrs['kl_divergence']:.4f}"
|
||||||
),
|
),
|
||||||
value=trial,
|
value=trial,
|
||||||
)
|
)
|
||||||
@@ -748,19 +817,38 @@ def run():
|
|||||||
print()
|
print()
|
||||||
print(f"Restoring model from trial [bold]{trial.user_attrs['index']}[/]...")
|
print(f"Restoring model from trial [bold]{trial.user_attrs['index']}[/]...")
|
||||||
print("* Parameters:")
|
print("* Parameters:")
|
||||||
for name, value in get_trial_parameters(trial).items():
|
for name, value in get_trial_parameters(settings, trial).items():
|
||||||
print(f" * {name} = [bold]{value}[/]")
|
print(f" * {name} = [bold]{value}[/]")
|
||||||
print("* Resetting model...")
|
if settings.use_ara_lora:
|
||||||
model.reset_model()
|
print("* Resetting model...")
|
||||||
print("* Abliterating...")
|
model.reset_model()
|
||||||
model.abliterate(
|
print("* Abliterating (Arbitrary-Rank Ablation with LoRA)...")
|
||||||
refusal_directions,
|
model.ara_lora_abliterate(
|
||||||
trial.user_attrs["direction_index"],
|
good_module_io,
|
||||||
{
|
bad_module_io,
|
||||||
k: AbliterationParameters(**v)
|
ARAParameters(**trial.user_attrs["ara_parameters"]),
|
||||||
for k, v in trial.user_attrs["parameters"].items()
|
)
|
||||||
},
|
elif settings.use_ara:
|
||||||
)
|
print("* Reloading model...")
|
||||||
|
model.reset_model()
|
||||||
|
print("* Abliterating (Arbitrary-Rank Ablation)...")
|
||||||
|
model.ara_abliterate(
|
||||||
|
good_module_io,
|
||||||
|
bad_module_io,
|
||||||
|
ARAParameters(**trial.user_attrs["ara_parameters"]),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
print("* Resetting model...")
|
||||||
|
model.reset_model()
|
||||||
|
print("* Abliterating...")
|
||||||
|
model.abliterate(
|
||||||
|
refusal_directions,
|
||||||
|
trial.user_attrs["direction_index"],
|
||||||
|
{
|
||||||
|
k: AbliterationParameters(**v)
|
||||||
|
for k, v in trial.user_attrs["parameters"].items()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
print()
|
print()
|
||||||
@@ -796,8 +884,12 @@ def run():
|
|||||||
print("Saving LoRA adapter...")
|
print("Saving LoRA adapter...")
|
||||||
model.model.save_pretrained(save_directory)
|
model.model.save_pretrained(save_directory)
|
||||||
else:
|
else:
|
||||||
print("Saving merged model...")
|
if settings.use_ara:
|
||||||
merged_model = model.get_merged_model()
|
print("Saving model...")
|
||||||
|
merged_model = model.model
|
||||||
|
else:
|
||||||
|
print("Saving merged model...")
|
||||||
|
merged_model = model.get_merged_model()
|
||||||
merged_model.save_pretrained(save_directory)
|
merged_model.save_pretrained(save_directory)
|
||||||
del merged_model
|
del merged_model
|
||||||
empty_cache()
|
empty_cache()
|
||||||
@@ -851,8 +943,12 @@ def run():
|
|||||||
token=token,
|
token=token,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
print("Uploading merged model...")
|
if settings.use_ara:
|
||||||
merged_model = model.get_merged_model()
|
print("Uploading model...")
|
||||||
|
merged_model = model.model
|
||||||
|
else:
|
||||||
|
print("Uploading merged model...")
|
||||||
|
merged_model = model.get_merged_model()
|
||||||
merged_model.push_to_hub(
|
merged_model.push_to_hub(
|
||||||
repo_id,
|
repo_id,
|
||||||
private=private,
|
private=private,
|
||||||
@@ -891,6 +987,14 @@ def run():
|
|||||||
card.data.tags.append("uncensored")
|
card.data.tags.append("uncensored")
|
||||||
card.data.tags.append("decensored")
|
card.data.tags.append("decensored")
|
||||||
card.data.tags.append("abliterated")
|
card.data.tags.append("abliterated")
|
||||||
|
if settings.use_ara:
|
||||||
|
card.data.tags.append("ara")
|
||||||
|
elif (
|
||||||
|
settings.orthogonalize_direction
|
||||||
|
and settings.row_normalization
|
||||||
|
== RowNormalization.FULL
|
||||||
|
):
|
||||||
|
card.data.tags.append("mpoa")
|
||||||
card.text = (
|
card.text = (
|
||||||
get_readme_intro(
|
get_readme_intro(
|
||||||
settings,
|
settings,
|
||||||
@@ -966,6 +1070,7 @@ def run():
|
|||||||
hflm = HFLM(
|
hflm = HFLM(
|
||||||
pretrained=model.model, # ty:ignore[invalid-argument-type]
|
pretrained=model.model, # ty:ignore[invalid-argument-type]
|
||||||
tokenizer=model.tokenizer, # ty:ignore[invalid-argument-type]
|
tokenizer=model.tokenizer, # ty:ignore[invalid-argument-type]
|
||||||
|
batch_size="auto",
|
||||||
)
|
)
|
||||||
|
|
||||||
table = Table()
|
table = Table()
|
||||||
@@ -989,7 +1094,6 @@ def run():
|
|||||||
results = lm_eval.simple_evaluate(
|
results = lm_eval.simple_evaluate(
|
||||||
model=hflm,
|
model=hflm,
|
||||||
tasks=[benchmark.task],
|
tasks=[benchmark.task],
|
||||||
batch_size="auto",
|
|
||||||
)
|
)
|
||||||
return results["results"][benchmark.task]
|
return results["results"][benchmark.task]
|
||||||
|
|
||||||
|
|||||||
+398
-16
@@ -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, 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,8 @@ 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 transformers import (
|
from transformers import (
|
||||||
AutoModelForCausalLM,
|
AutoModelForCausalLM,
|
||||||
AutoModelForImageTextToText,
|
AutoModelForImageTextToText,
|
||||||
@@ -30,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(
|
||||||
@@ -52,6 +54,23 @@ class AbliterationParameters:
|
|||||||
min_weight_distance: float
|
min_weight_distance: float
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ARAParameters:
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
# 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
|
||||||
@@ -142,7 +161,8 @@ 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()
|
if not settings.use_ara or settings.use_ara_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.
|
||||||
@@ -168,21 +188,24 @@ class Model:
|
|||||||
# because hybrid models like Qwen3.5 MoE have modules with different names
|
# because hybrid models like Qwen3.5 MoE have modules with different names
|
||||||
# across layers (e.g. "o_proj" on attention layers, "out_proj" on linear attention layers).
|
# across layers (e.g. "o_proj" on attention layers, "out_proj" on linear attention layers).
|
||||||
target_modules_set: set[str] = set()
|
target_modules_set: set[str] = set()
|
||||||
|
|
||||||
|
module_id_to_full_name = {
|
||||||
|
id(module): module_name
|
||||||
|
for module_name, module in self.model.named_modules()
|
||||||
|
}
|
||||||
|
|
||||||
for layer_index, layer in enumerate(self.get_layers()):
|
for layer_index in range(len(self.get_layers())):
|
||||||
module_id_to_leaf_name = {
|
|
||||||
id(module): module_name.split(".")[-1]
|
|
||||||
for module_name, module in layer.named_modules()
|
|
||||||
}
|
|
||||||
|
|
||||||
for modules in self.get_layer_modules(layer_index).values():
|
for modules in self.get_layer_modules(layer_index).values():
|
||||||
for module in modules:
|
for module in modules:
|
||||||
if id(module) in module_id_to_leaf_name:
|
full_name = module_id_to_full_name.get(id(module))
|
||||||
target_modules_set.add(module_id_to_leaf_name[id(module)])
|
if full_name is not None:
|
||||||
|
target_modules_set.add(full_name)
|
||||||
|
|
||||||
target_modules = list(target_modules_set)
|
target_modules = sorted(target_modules_set)
|
||||||
|
|
||||||
if self.settings.row_normalization != RowNormalization.FULL:
|
if self.settings.use_ara_lora:
|
||||||
|
lora_rank = self.settings.ara_lora_rank
|
||||||
|
elif self.settings.row_normalization != RowNormalization.FULL:
|
||||||
# Rank 1 is sufficient for directional ablation without renormalization.
|
# Rank 1 is sufficient for directional ablation without renormalization.
|
||||||
lora_rank = 1
|
lora_rank = 1
|
||||||
else:
|
else:
|
||||||
@@ -204,7 +227,10 @@ class Model:
|
|||||||
# so the result is a PeftModel rather than a PeftMixedModel.
|
# so the result is a PeftModel rather than a PeftMixedModel.
|
||||||
self.model = cast(PeftModel, get_peft_model(self.model, self.peft_config))
|
self.model = cast(PeftModel, get_peft_model(self.model, self.peft_config))
|
||||||
|
|
||||||
print(f"* LoRA adapters initialized (targets: {', '.join(target_modules)})")
|
display_targets = sorted({name.rsplit(".", 1)[-1] for name in target_modules})
|
||||||
|
print(
|
||||||
|
f"* LoRA adapters initialized (target types: {', '.join(display_targets)})"
|
||||||
|
)
|
||||||
|
|
||||||
def _get_quantization_config(self, dtype: str) -> BitsAndBytesConfig | None:
|
def _get_quantization_config(self, dtype: str) -> BitsAndBytesConfig | None:
|
||||||
"""
|
"""
|
||||||
@@ -288,7 +314,11 @@ class Model:
|
|||||||
performs full model reload with quantization config.
|
performs full model reload with quantization config.
|
||||||
"""
|
"""
|
||||||
current_model = getattr(self.model.config, "name_or_path", None)
|
current_model = getattr(self.model.config, "name_or_path", None)
|
||||||
if current_model == self.settings.model and not self.needs_reload:
|
if (
|
||||||
|
current_model == self.settings.model
|
||||||
|
and not self.needs_reload
|
||||||
|
and (not self.settings.use_ara or self.settings.use_ara_lora)
|
||||||
|
):
|
||||||
# Reset LoRA adapters to zero (identity transformation)
|
# Reset LoRA adapters to zero (identity transformation)
|
||||||
for name, module in self.model.named_modules():
|
for name, module in self.model.named_modules():
|
||||||
if "lora_B" in name and hasattr(module, "weight"):
|
if "lora_B" in name and hasattr(module, "weight"):
|
||||||
@@ -317,7 +347,8 @@ class Model:
|
|||||||
**extra_kwargs,
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._apply_lora()
|
if not self.settings.use_ara or self.settings.use_ara_lora:
|
||||||
|
self._apply_lora()
|
||||||
|
|
||||||
self.needs_reload = False
|
self.needs_reload = False
|
||||||
|
|
||||||
@@ -341,6 +372,9 @@ class Model:
|
|||||||
modules = {}
|
modules = {}
|
||||||
|
|
||||||
def try_add(component: str, module: Any):
|
def try_add(component: str, module: Any):
|
||||||
|
if component not in self.settings.target_components:
|
||||||
|
return
|
||||||
|
|
||||||
# Only add if it's a proper nn.Module (PEFT can wrap these with LoRA)
|
# Only add if it's a proper nn.Module (PEFT can wrap these with LoRA)
|
||||||
if isinstance(module, Module):
|
if isinstance(module, Module):
|
||||||
if component not in modules:
|
if component not in modules:
|
||||||
@@ -541,6 +575,228 @@ 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,
|
||||||
|
parameters: ARAParameters,
|
||||||
|
):
|
||||||
|
for layer_index in range(
|
||||||
|
parameters.start_layer_index,
|
||||||
|
parameters.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
|
||||||
|
|
||||||
|
row_norms = LA.vector_norm(matrix, dim=1, keepdim=True).detach()
|
||||||
|
|
||||||
|
# Helper function for reparameterization (row-norm preservation constraint).
|
||||||
|
def get_matrix() -> Tensor:
|
||||||
|
if self.settings.row_normalization == RowNormalization.FULL:
|
||||||
|
# See https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration
|
||||||
|
return row_norms * F.normalize(matrix, p=2, dim=1)
|
||||||
|
else:
|
||||||
|
return matrix
|
||||||
|
|
||||||
|
good_input, good_output = good_module_io[layer_index][component][
|
||||||
|
module_index
|
||||||
|
]
|
||||||
|
bad_input, bad_output = bad_module_io[layer_index][component][
|
||||||
|
module_index
|
||||||
|
]
|
||||||
|
|
||||||
|
good_input = good_input.to(matrix.device)
|
||||||
|
good_output = good_output.to(matrix.device)
|
||||||
|
bad_input = bad_input.to(matrix.device)
|
||||||
|
bad_output = bad_output.to(matrix.device)
|
||||||
|
|
||||||
|
def objective(matrix: Tensor) -> Tensor:
|
||||||
|
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 = (
|
||||||
|
(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
|
||||||
|
)
|
||||||
|
|
||||||
|
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(get_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}"
|
||||||
|
# )
|
||||||
|
|
||||||
|
# Free the gradient buffers accumulated on the weight parameters
|
||||||
|
# during optimization. Without this, they persist on the model
|
||||||
|
# (one full-size gradient per processed weight) and can easily
|
||||||
|
# consume tens of GiB of VRAM, causing out-of-memory errors
|
||||||
|
# during the subsequent evaluation.
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
matrix.copy_(get_matrix())
|
||||||
|
|
||||||
|
def ara_lora_abliterate(
|
||||||
|
self,
|
||||||
|
good_module_io: ModuleIO,
|
||||||
|
bad_module_io: ModuleIO,
|
||||||
|
parameters: ARAParameters,
|
||||||
|
):
|
||||||
|
for layer_index in range(
|
||||||
|
parameters.start_layer_index,
|
||||||
|
parameters.end_layer_index,
|
||||||
|
):
|
||||||
|
for component, modules in self.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)
|
||||||
|
|
||||||
|
# Base weight handling and dequantization.
|
||||||
|
# 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:
|
||||||
|
# Maintain the original dequantization logic for bitsandbytes.
|
||||||
|
W_base = cast(
|
||||||
|
Tensor,
|
||||||
|
bnb.functional.dequantize_4bit(
|
||||||
|
base_weight.data,
|
||||||
|
quant_state
|
||||||
|
).to(torch.float32),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Row normalization setup.
|
||||||
|
# Pre-calculate the original row norms to preserve them.
|
||||||
|
# This implements the RowNormalization.FULL logic.
|
||||||
|
W_row_norms = LA.vector_norm(W_base, dim=1, keepdim=True).detach()
|
||||||
|
|
||||||
|
# Adapter target identification.
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
# Data preparation.
|
||||||
|
# Move I/O tensors to the device of the adapter weights.
|
||||||
|
good_input, good_output = good_module_io[layer_index][component][module_index]
|
||||||
|
bad_input, bad_output = 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)
|
||||||
|
|
||||||
|
# The objective function.
|
||||||
|
def objective(A: Tensor, B: Tensor) -> Tensor:
|
||||||
|
# Calculate effective weight: W_eff = W_base + B @ A.
|
||||||
|
W_eff = W_base + (B @ A)
|
||||||
|
|
||||||
|
# Apply Row Normalization (keep original norms).
|
||||||
|
if self.settings.row_normalization == RowNormalization.FULL:
|
||||||
|
# Normalize to unit length, then scale by original 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
|
||||||
|
|
||||||
|
# The original ARA loss function.
|
||||||
|
preserve_good_behavior = (
|
||||||
|
(new_good_output - good_output) ** 2
|
||||||
|
).mean()
|
||||||
|
|
||||||
|
steer_bad_behavior = (
|
||||||
|
mean_distances_to_knn(
|
||||||
|
new_bad_output,
|
||||||
|
good_output,
|
||||||
|
parameters.neighbor_count,
|
||||||
|
).mean()
|
||||||
|
+ 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
|
||||||
|
)
|
||||||
|
|
||||||
|
# Optimization loop.
|
||||||
|
# We optimize A and B, not the base matrix.
|
||||||
|
optimizer = LBFGS(
|
||||||
|
[lora_A, lora_B],
|
||||||
|
lr=1.0,
|
||||||
|
max_iter=20,
|
||||||
|
history_size=10,
|
||||||
|
line_search_fn="strong_wolfe",
|
||||||
|
)
|
||||||
|
|
||||||
|
def closure():
|
||||||
|
optimizer.zero_grad()
|
||||||
|
# Pass the actual tensors being optimized to the objective.
|
||||||
|
loss = objective(lora_A, lora_B)
|
||||||
|
loss.backward()
|
||||||
|
return loss
|
||||||
|
|
||||||
|
# Run optimization steps.
|
||||||
|
for step in range(5):
|
||||||
|
optimizer.step(closure)
|
||||||
|
|
||||||
|
# Free the gradient buffers accumulated on the LoRA adapter
|
||||||
|
# parameters during optimization (see ara_abliterate for details).
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
def generate(
|
def generate(
|
||||||
self,
|
self,
|
||||||
prompts: list[Prompt],
|
prompts: list[Prompt],
|
||||||
@@ -675,6 +931,132 @@ class Model:
|
|||||||
|
|
||||||
return torch.cat(residuals, dim=0)
|
return torch.cat(residuals, dim=0)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
# We work with logprobs rather than probabilities for numerical stability
|
# We work with logprobs rather than probabilities for numerical stability
|
||||||
# when computing the KL divergence.
|
# when computing the KL divergence.
|
||||||
def get_logprobs(self, prompts: list[Prompt]) -> Tensor:
|
def get_logprobs(self, prompts: list[Prompt]) -> Tensor:
|
||||||
|
|||||||
+55
-14
@@ -25,8 +25,9 @@ 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, RowNormalization, Settings
|
||||||
|
|
||||||
print = Console(highlight=False).print
|
print = Console(highlight=False).print
|
||||||
|
|
||||||
@@ -234,6 +235,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.
|
||||||
@@ -256,19 +265,46 @@ def empty_cache():
|
|||||||
gc.collect()
|
gc.collect()
|
||||||
|
|
||||||
|
|
||||||
def get_trial_parameters(trial: Trial) -> dict[str, str]:
|
def get_trial_parameters(settings: Settings, trial: Trial) -> dict[str, str]:
|
||||||
params = {}
|
if settings.use_ara:
|
||||||
|
parameters = trial.user_attrs["ara_parameters"]
|
||||||
|
|
||||||
direction_index = trial.user_attrs["direction_index"]
|
return {
|
||||||
params["direction_index"] = (
|
name: (f"{value:.4f}" if isinstance(value, float) else f"{value}")
|
||||||
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
|
for name, value in parameters.items()
|
||||||
)
|
}
|
||||||
|
else:
|
||||||
|
params = {}
|
||||||
|
|
||||||
for component, parameters in trial.user_attrs["parameters"].items():
|
direction_index = trial.user_attrs["direction_index"]
|
||||||
for name, value in parameters.items():
|
params["direction_index"] = (
|
||||||
params[f"{component}.{name}"] = f"{value:.2f}"
|
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
|
||||||
|
)
|
||||||
|
|
||||||
return params
|
for component, parameters in trial.user_attrs["parameters"].items():
|
||||||
|
for name, value in parameters.items():
|
||||||
|
params[f"{component}.{name}"] = f"{value:.2f}"
|
||||||
|
|
||||||
|
return params
|
||||||
|
|
||||||
|
|
||||||
|
def get_method_description(settings: Settings) -> str:
|
||||||
|
if settings.use_ara:
|
||||||
|
return (
|
||||||
|
" with the [Arbitrary-Rank Ablation (ARA)](https://github.com/p-e-w/heretic/pull/211) method"
|
||||||
|
+ (
|
||||||
|
" (with row-norm preservation)"
|
||||||
|
if settings.row_normalization == RowNormalization.FULL
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
)
|
||||||
|
elif (
|
||||||
|
settings.orthogonalize_direction
|
||||||
|
and settings.row_normalization == RowNormalization.FULL
|
||||||
|
):
|
||||||
|
return " with a variant of the [Magnitude-Preserving Orthogonal Ablation (MPOA)](https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration) method"
|
||||||
|
else:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
def get_readme_intro(
|
def get_readme_intro(
|
||||||
@@ -285,7 +321,9 @@ def get_readme_intro(
|
|||||||
|
|
||||||
return f"""# This is a decensored version of {
|
return f"""# This is a decensored version of {
|
||||||
model_link
|
model_link
|
||||||
}, made using [Heretic](https://github.com/p-e-w/heretic) v{version("heretic-llm")}
|
}, made using [Heretic](https://github.com/p-e-w/heretic) v{version("heretic-llm")}{
|
||||||
|
get_method_description(settings)
|
||||||
|
}
|
||||||
|
|
||||||
## Abliteration parameters
|
## Abliteration parameters
|
||||||
|
|
||||||
@@ -295,7 +333,7 @@ def get_readme_intro(
|
|||||||
chr(10).join(
|
chr(10).join(
|
||||||
[
|
[
|
||||||
f"| **{name}** | {value} |"
|
f"| **{name}** | {value} |"
|
||||||
for name, value in get_trial_parameters(trial).items()
|
for name, value in get_trial_parameters(settings, trial).items()
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -304,7 +342,10 @@ def get_readme_intro(
|
|||||||
|
|
||||||
| Metric | This model | Original model ({model_link}) |
|
| Metric | This model | Original model ({model_link}) |
|
||||||
| :----- | :--------: | :---------------------------: |
|
| :----- | :--------: | :---------------------------: |
|
||||||
| **KL divergence** | {trial.user_attrs["kl_divergence"]:.4f} | 0 *(by definition)* |
|
| **{"PIQA acc_norm" if settings.use_piqa else "KL divergence"}** | {
|
||||||
|
(-1 if settings.use_piqa else 1) * trial.user_attrs["kl_divergence"]:.4f} | {
|
||||||
|
"*Unknown*" if settings.use_piqa else "0 *(by definition)*"
|
||||||
|
} |
|
||||||
| **Refusals** | {trial.user_attrs["refusals"]}/{len(bad_prompts)} | {base_refusals}/{
|
| **Refusals** | {trial.user_attrs["refusals"]}/{len(bad_prompts)} | {base_refusals}/{
|
||||||
len(bad_prompts)
|
len(bad_prompts)
|
||||||
} |
|
} |
|
||||||
|
|||||||
Reference in New Issue
Block a user