feat(ara): implement optimization for ARA parameters

This commit is contained in:
Philipp Emanuel Weidmann
2026-03-02 14:36:57 +05:30
parent 154241f8a2
commit b8f4a9c985
4 changed files with 237 additions and 153 deletions
+16
View File
@@ -176,6 +176,22 @@ class Settings(BaseSettings):
), ),
) )
target_components: list[str] = Field(
default=["attn.o_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."
),
)
orthogonalize_direction: bool = Field( orthogonalize_direction: bool = Field(
default=False, default=False,
description=( description=(
+72 -30
View File
@@ -206,8 +206,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.[/]"
) )
if not settings.use_ara:
# 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)
@@ -411,37 +412,13 @@ def run():
evaluator.get_score() evaluator.get_score()
return return
def tensor_shape_repr(self: torch.Tensor): if settings.use_ara:
return f"tensor(shape={tuple(self.shape)}, dtype={self.dtype}, device={self.device})"
torch.Tensor.__repr__ = tensor_shape_repr # ty:ignore[invalid-assignment]
print() print()
print("Obtaining module I/O for good prompts...") print("Obtaining module I/O for good prompts...")
good_module_io = model.get_module_io_batched(good_prompts) good_module_io = model.get_module_io_batched(good_prompts)
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)
else:
# 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...")
print("* Obtaining residuals for good prompts...") print("* Obtaining residuals for good prompts...")
@@ -486,6 +463,33 @@ def run():
trial_index += 1 trial_index += 1
trial.set_user_attr("index", trial_index) trial.set_user_attr("index", trial_index)
if settings.use_ara:
start_layer_index = trial.suggest_int(
"start_layer_index",
0,
len(model.get_layers()) // 3,
)
end_layer_index = trial.suggest_int(
"end_layer_index",
len(model.get_layers()) // 2,
len(model.get_layers()),
)
preserve_good_behavior_weight = trial.suggest_float(
"preserve_good_behavior_weight",
0.0,
1.0,
)
steer_bad_behavior_weight = trial.suggest_float(
"steer_bad_behavior_weight",
0.0,
1.0,
)
tie_to_original_matrix_weight = trial.suggest_float(
"tie_to_original_matrix_weight",
0.2, # Minimum to prevent "optimizing" away the regularization term.
1.0,
)
else:
direction_scope = trial.suggest_categorical( direction_scope = trial.suggest_categorical(
"direction_scope", "direction_scope",
[ [
@@ -550,15 +554,31 @@ def run():
) )
trial.set_user_attr("direction_index", direction_index) trial.set_user_attr("direction_index", direction_index)
trial.set_user_attr("parameters", {k: asdict(v) for k, v in parameters.items()}) 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}[/]")
if settings.use_ara:
print("* Reloading model...")
model.reset_model()
print("* Abliterating (Arbitrary-Rank Ablation)...")
model.ara_abliterate(
good_module_io,
bad_module_io,
start_layer_index,
end_layer_index,
preserve_good_behavior_weight,
steer_bad_behavior_weight,
tie_to_original_matrix_weight,
)
else:
print("* Resetting model...") print("* Resetting model...")
model.reset_model() model.reset_model()
print("* Abliterating...") print("* Abliterating...")
@@ -739,8 +759,22 @@ 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}[/]")
if settings.use_ara:
print("* Reloading model...")
model.reset_model()
print("* Abliterating (Arbitrary-Rank Ablation)...")
model.ara_abliterate(
good_module_io,
bad_module_io,
trial.params["start_layer_index"],
trial.params["end_layer_index"],
trial.params["preserve_good_behavior_weight"],
trial.params["steer_bad_behavior_weight"],
trial.params["tie_to_original_matrix_weight"],
)
else:
print("* Resetting model...") print("* Resetting model...")
model.reset_model() model.reset_model()
print("* Abliterating...") print("* Abliterating...")
@@ -785,6 +819,10 @@ def run():
if strategy == "adapter": if strategy == "adapter":
print("Saving LoRA adapter...") print("Saving LoRA adapter...")
model.model.save_pretrained(save_directory) model.model.save_pretrained(save_directory)
else:
if settings.use_ara:
print("Saving model...")
merged_model = model.model
else: else:
print("Saving merged model...") print("Saving merged model...")
merged_model = model.get_merged_model() merged_model = model.get_merged_model()
@@ -838,6 +876,10 @@ def run():
private=private, private=private,
token=token, token=token,
) )
else:
if settings.use_ara:
print("Uploading model...")
merged_model = model.model
else: else:
print("Uploading merged model...") print("Uploading merged model...")
merged_model = model.get_merged_model() merged_model = model.get_merged_model()
+22 -5
View File
@@ -153,7 +153,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:
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.
@@ -285,7 +286,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
):
# 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"):
@@ -314,6 +319,7 @@ class Model:
**extra_kwargs, **extra_kwargs,
) )
if not self.settings.use_ara:
self._apply_lora() self._apply_lora()
self.needs_reload = False self.needs_reload = False
@@ -338,6 +344,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:
@@ -564,6 +573,14 @@ class Model:
# On average, the outputs for "bad" prompts should resemble # On average, the outputs for "bad" prompts should resemble
# the original outputs for "good" prompts (which steers the # the original outputs for "good" prompts (which steers the
# behavior for "bad" prompts towards that for "good" prompts). # 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.
steer_bad_behavior = ( steer_bad_behavior = (
( (
(bad_input @ matrix.T).mean(dim=0) (bad_input @ matrix.T).mean(dim=0)
@@ -602,9 +619,9 @@ class Model:
# Convergence usually happens within 2-3 steps, so this is more than enough. # Convergence usually happens within 2-3 steps, so this is more than enough.
for step in range(5): for step in range(5):
loss = optimizer.step(closure) loss = optimizer.step(closure)
print( # print(
f"\\[{layer_index}/{component}/{module_index}] Step: {step}, Loss: {loss.item():.6f}" # f"\\[{layer_index}/{component}/{module_index}] Step: {step}, Loss: {loss.item():.6f}"
) # )
def generate( def generate(
self, self,
+11 -2
View File
@@ -250,7 +250,16 @@ 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]:
if settings.use_ara:
return {
"start_layer_index": f"{trial.params['start_layer_index']}",
"end_layer_index": f"{trial.params['end_layer_index']}",
"preserve_good_behavior_weight": f"{trial.params['preserve_good_behavior_weight']:.4f}",
"steer_bad_behavior_weight": f"{trial.params['steer_bad_behavior_weight']:.4f}",
"tie_to_original_matrix_weight": f"{trial.params['tie_to_original_matrix_weight']:.4f}",
}
else:
params = {} params = {}
direction_index = trial.user_attrs["direction_index"] direction_index = trial.user_attrs["direction_index"]
@@ -285,7 +294,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()
] ]
) )
} }