mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-29 07:21:27 -07:00
feat(ara): implement optimization for ARA parameters
This commit is contained in:
@@ -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=(
|
||||||
|
|||||||
+178
-136
@@ -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.[/]"
|
||||||
)
|
)
|
||||||
|
|
||||||
# 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)
|
||||||
@@ -411,71 +412,47 @@ 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})"
|
print()
|
||||||
|
print("Obtaining module I/O for good prompts...")
|
||||||
|
good_module_io = model.get_module_io_batched(good_prompts)
|
||||||
|
print("Obtaining module I/O for 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)
|
||||||
|
|
||||||
torch.Tensor.__repr__ = tensor_shape_repr # ty:ignore[invalid-assignment]
|
good_means = good_residuals.mean(dim=0)
|
||||||
|
bad_means = bad_residuals.mean(dim=0)
|
||||||
|
|
||||||
print()
|
refusal_directions = F.normalize(bad_means - good_means, p=2, dim=1)
|
||||||
print("Obtaining module I/O for good prompts...")
|
|
||||||
good_module_io = model.get_module_io_batched(good_prompts)
|
|
||||||
print("Obtaining module I/O for bad prompts...")
|
|
||||||
bad_module_io = model.get_module_io_batched(bad_prompts)
|
|
||||||
|
|
||||||
# print(good_module_io)
|
if settings.orthogonalize_direction:
|
||||||
|
# Implements https://huggingface.co/blog/grimjim/projected-abliteration
|
||||||
|
# Adjust the refusal directions so that only the component that is
|
||||||
|
# orthogonal to the good direction is subtracted during abliteration.
|
||||||
|
good_directions = F.normalize(good_means, p=2, dim=1)
|
||||||
|
projection_vector = torch.sum(refusal_directions * good_directions, dim=1)
|
||||||
|
refusal_directions = (
|
||||||
|
refusal_directions - projection_vector.unsqueeze(1) * good_directions
|
||||||
|
)
|
||||||
|
refusal_directions = F.normalize(refusal_directions, p=2, dim=1)
|
||||||
|
|
||||||
print()
|
analyzer = Analyzer(settings, model, good_residuals, bad_residuals)
|
||||||
print("Performing Arbitrary-Rank Ablation...")
|
|
||||||
|
|
||||||
model.ara_abliterate(
|
if settings.print_residual_geometry:
|
||||||
good_module_io,
|
analyzer.print_residual_geometry()
|
||||||
bad_module_io,
|
|
||||||
0,
|
|
||||||
len(model.get_layers()),
|
|
||||||
1.0,
|
|
||||||
1.0,
|
|
||||||
1.0,
|
|
||||||
)
|
|
||||||
|
|
||||||
print()
|
if settings.plot_residuals:
|
||||||
print("Evaluating...")
|
analyzer.plot_residuals()
|
||||||
evaluator.get_score()
|
|
||||||
return
|
|
||||||
|
|
||||||
print()
|
# We don't need the residuals after computing refusal directions.
|
||||||
print("Calculating per-layer refusal directions...")
|
del good_residuals, bad_residuals, analyzer
|
||||||
print("* Obtaining residuals for good prompts...")
|
empty_cache()
|
||||||
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)
|
|
||||||
bad_means = bad_residuals.mean(dim=0)
|
|
||||||
|
|
||||||
refusal_directions = F.normalize(bad_means - good_means, p=2, dim=1)
|
|
||||||
|
|
||||||
if settings.orthogonalize_direction:
|
|
||||||
# Implements https://huggingface.co/blog/grimjim/projected-abliteration
|
|
||||||
# Adjust the refusal directions so that only the component that is
|
|
||||||
# orthogonal to the good direction is subtracted during abliteration.
|
|
||||||
good_directions = F.normalize(good_means, p=2, dim=1)
|
|
||||||
projection_vector = torch.sum(refusal_directions * good_directions, dim=1)
|
|
||||||
refusal_directions = (
|
|
||||||
refusal_directions - projection_vector.unsqueeze(1) * good_directions
|
|
||||||
)
|
|
||||||
refusal_directions = F.normalize(refusal_directions, p=2, dim=1)
|
|
||||||
|
|
||||||
analyzer = Analyzer(settings, model, good_residuals, bad_residuals)
|
|
||||||
|
|
||||||
if settings.print_residual_geometry:
|
|
||||||
analyzer.print_residual_geometry()
|
|
||||||
|
|
||||||
if settings.plot_residuals:
|
|
||||||
analyzer.plot_residuals()
|
|
||||||
|
|
||||||
# We don't need the residuals after computing refusal directions.
|
|
||||||
del good_residuals, bad_residuals, analyzer
|
|
||||||
empty_cache()
|
|
||||||
|
|
||||||
trial_index = 0
|
trial_index = 0
|
||||||
start_index = 0
|
start_index = 0
|
||||||
@@ -486,83 +463,126 @@ 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()) // 3,
|
||||||
],
|
|
||||||
)
|
|
||||||
|
|
||||||
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.0,
|
||||||
1.0,
|
1.0,
|
||||||
0.6 * last_layer_index,
|
)
|
||||||
|
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",
|
||||||
|
[
|
||||||
|
"global",
|
||||||
|
"per layer",
|
||||||
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
parameters[component] = AbliterationParameters(
|
last_layer_index = len(model.get_layers()) - 1
|
||||||
max_weight=max_weight,
|
|
||||||
max_weight_position=max_weight_position,
|
# Discrimination between "harmful" and "harmless" inputs is usually strongest
|
||||||
min_weight=(min_weight * max_weight),
|
# in layers slightly past the midpoint of the layer stack. See the original
|
||||||
min_weight_distance=min_weight_distance,
|
# 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,
|
||||||
)
|
)
|
||||||
|
|
||||||
trial.set_user_attr("direction_index", direction_index)
|
if direction_scope == "per layer":
|
||||||
trial.set_user_attr("parameters", {k: asdict(v) for k, v in parameters.items()})
|
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:
|
||||||
model.reset_model()
|
print("* Reloading model...")
|
||||||
print("* Abliterating...")
|
model.reset_model()
|
||||||
model.abliterate(refusal_directions, direction_index, parameters)
|
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...")
|
||||||
|
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()
|
||||||
|
|
||||||
@@ -739,19 +759,33 @@ 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:
|
||||||
model.reset_model()
|
print("* Reloading model...")
|
||||||
print("* Abliterating...")
|
model.reset_model()
|
||||||
model.abliterate(
|
print("* Abliterating (Arbitrary-Rank Ablation)...")
|
||||||
refusal_directions,
|
model.ara_abliterate(
|
||||||
trial.user_attrs["direction_index"],
|
good_module_io,
|
||||||
{
|
bad_module_io,
|
||||||
k: AbliterationParameters(**v)
|
trial.params["start_layer_index"],
|
||||||
for k, v in trial.user_attrs["parameters"].items()
|
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...")
|
||||||
|
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()
|
||||||
@@ -786,8 +820,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()
|
||||||
@@ -839,8 +877,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,
|
||||||
|
|||||||
+23
-6
@@ -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,7 +319,8 @@ class Model:
|
|||||||
**extra_kwargs,
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._apply_lora()
|
if not self.settings.use_ara:
|
||||||
|
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,
|
||||||
|
|||||||
+20
-11
@@ -250,19 +250,28 @@ 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:
|
||||||
|
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 = {}
|
||||||
|
|
||||||
direction_index = trial.user_attrs["direction_index"]
|
direction_index = trial.user_attrs["direction_index"]
|
||||||
params["direction_index"] = (
|
params["direction_index"] = (
|
||||||
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
|
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
|
||||||
)
|
)
|
||||||
|
|
||||||
for component, parameters in trial.user_attrs["parameters"].items():
|
for component, parameters in trial.user_attrs["parameters"].items():
|
||||||
for name, value in parameters.items():
|
for name, value in parameters.items():
|
||||||
params[f"{component}.{name}"] = f"{value:.2f}"
|
params[f"{component}.{name}"] = f"{value:.2f}"
|
||||||
|
|
||||||
return params
|
return params
|
||||||
|
|
||||||
|
|
||||||
def get_readme_intro(
|
def get_readme_intro(
|
||||||
@@ -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()
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user