feat: re-implement ARA as a modifier plugin

Arbitrary-Rank Ablation (ARA) (Weidmann 2026) was originally introduced by @p-e-w in #211

Co-authored-by: kabachuha <artemkhrapov2001@yandex.ru>
Co-authored-by: joninco <joninco@bullpoint.org>
Co-authored-by: Ashar <coder3101@users.noreply.github.com>
This commit is contained in:
Philipp Emanuel Weidmann
2026-09-28 19:24:08 +05:30
co-authored by kabachuha joninco Ashar
parent ffa66af2d4
commit ce8b05c77b
11 changed files with 575 additions and 11 deletions
+55 -1
View File
@@ -84,7 +84,7 @@ scorers = [
# Note that only a single modifier can currently be applied, # Note that only a single modifier can currently be applied,
# and this list must contain exactly one entry. # and this list must contain exactly one entry.
modifiers = [ modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" }, { plugin = "heretic.modifiers.ara.ARA" },
] ]
# Number of abliteration trials to run during optimization. # Number of abliteration trials to run during optimization.
@@ -119,6 +119,7 @@ dataset = "mlabonne/harmful_behaviors"
split = "train[:100]" split = "train[:100]"
column = "text" column = "text"
# Plugin-specific settings live in top-level TOML tables. # Plugin-specific settings live in top-level TOML tables.
# #
# For scorer plugins, use: `[scorer.<ClassName>]` (and optionally `[scorer.<ClassName>_<instance_name>]` for instance-related config). # For scorer plugins, use: `[scorer.<ClassName>]` (and optionally `[scorer.<ClassName>_<instance_name>]` for instance-related config).
@@ -198,12 +199,65 @@ dataset = "mlabonne/harmful_behaviors"
split = "test[:100]" split = "test[:100]"
column = "text" column = "text"
# Dataset of prompts used to measure KL divergence from original model. # Dataset of prompts used to measure KL divergence from original model.
[scorer.KLDivergence.prompts] [scorer.KLDivergence.prompts]
dataset = "mlabonne/harmless_alpaca" dataset = "mlabonne/harmless_alpaca"
split = "test[:100]" split = "test[:100]"
column = "text" column = "text"
[scorer.BenchmarkScore]
# Name that describes what the configured benchmark score measures.
score_name = "PIQA acc_norm"
# Task ID of the benchmark in the Language Model Evaluation Harness.
task = "piqa"
# Task metric to use as the benchmark score.
metric = "acc_norm,none"
[modifier.ARA]
# Whether to renormalize the rows of the modified matrices to preserve
# the original matrices' row magnitudes. This is believed to improve
# intelligence retention (see Lai 2025, "Magnitude-Preserving Orthogonal Ablation").
preserve_row_magnitudes = true
# The rank of the LoRA adapter to use.
# While mathematically, ARA is of "arbitrary" rank, experiments have shown that
# singular values tend to drop rapidly after a few dozen dimensions, and approximating
# the full transformation with a LoRA has many practical advantages.
lora_rank = 50
# Number of (outer) L-BFGS optimization steps to perform.
n_optimization_steps = 5
# Learning rate to use in the L-BFGS optimizer.
learning_rate = 1.0
# Maximum number of (inner) iterations to perform per (outer) L-BFGS optimization step.
max_iter = 20
# Number of past updates to store for approximating the Hessian matrix in the L-BFGS optimizer.
history_size = 10
# Whether to print the loss value for each L-BFGS optimization step.
print_loss = false
# Dataset of prompts that tend to produce desirable responses.
[modifier.ARA.good_prompts]
dataset = "mlabonne/harmless_alpaca"
split = "train[:400]"
column = "text"
# Dataset of prompts that tend to produce undesirable responses.
[modifier.ARA.bad_prompts]
dataset = "mlabonne/harmful_behaviors"
split = "train[:400]"
column = "text"
[modifier.Abliteration] [modifier.Abliteration]
# Whether to adjust the residual directions so that only the component that is # Whether to adjust the residual 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.
+1 -1
View File
@@ -392,7 +392,7 @@ class Settings(BaseSettings):
modifiers: list[ModifierConfig] = Field( modifiers: list[ModifierConfig] = Field(
default=[ default=[
ModifierConfig( ModifierConfig(
plugin="heretic.modifiers.abliteration.Abliteration", plugin="heretic.modifiers.ara.ARA",
), ),
], ],
description=( description=(
+2 -2
View File
@@ -630,7 +630,7 @@ def run():
print(f" * {name} = [bold]{value}[/]") print(f" * {name} = [bold]{value}[/]")
print("* Resetting model...") print("* Resetting model...")
modifier.reset_model(ctx) modifier.reset_model(ctx)
print(f"* Modifying model ({modifier_name})...") print(f"* Modifying model using {modifier_name}...")
modifier.modify_model(ctx, parameters) modifier.modify_model(ctx, parameters)
print("* Evaluating...") print("* Evaluating...")
scores = evaluator.get_scores() scores = evaluator.get_scores()
@@ -868,7 +868,7 @@ def run():
ctx = Context(settings=settings, model=model) ctx = Context(settings=settings, model=model)
print("* Resetting model...") print("* Resetting model...")
modifier.reset_model(ctx) modifier.reset_model(ctx)
print(f"* Modifying model ({modifier_name})...") print(f"* Modifying model using {modifier_name}...")
parameters = modifier.parameters_class.from_dict( parameters = modifier.parameters_class.from_dict(
trial.user_attrs["parameters"] trial.user_attrs["parameters"]
) )
+135 -1
View File
@@ -2,12 +2,13 @@
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors # Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
from contextlib import suppress from contextlib import suppress
from typing import Any, Type, cast from typing import Any, Callable, Type, TypeAlias, cast
import torch import torch
from peft import LoraConfig, PeftModel, get_peft_model from peft import LoraConfig, PeftModel, get_peft_model
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.utils.hooks import RemovableHandle
from transformers import ( from transformers import (
AutoModelForCausalLM, AutoModelForCausalLM,
AutoModelForImageTextToText, AutoModelForImageTextToText,
@@ -41,6 +42,13 @@ def get_model_class(
return AutoModelForCausalLM return AutoModelForCausalLM
# 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
@@ -616,6 +624,132 @@ class Model:
return (running_sum / total_count).to(torch.float32) return (running_sum / total_count).to(torch.float32)
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
def get_logits(self, prompts: list[Prompt]) -> Tensor: def get_logits(self, prompts: list[Prompt]) -> Tensor:
# We only generate one token, and we return the raw logits over the vocabulary # We only generate one token, and we return the raw logits over the vocabulary
# at that token position, for each prompt. # at that token position, for each prompt.
+6 -6
View File
@@ -160,27 +160,27 @@ class Abliteration(Modifier[Parameters]):
print( print(
f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..." f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..."
) )
self.good_prompts = ctx.load_prompts(self.settings.good_prompts) good_prompts = ctx.load_prompts(self.settings.good_prompts)
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded") print(f"* [bold]{len(good_prompts)}[/] prompts loaded")
print() print()
print( print(
f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..." f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..."
) )
self.bad_prompts = ctx.load_prompts(self.settings.bad_prompts) bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
print(f"* [bold]{len(self.bad_prompts)}[/] prompts loaded") print(f"* [bold]{len(bad_prompts)}[/] prompts loaded")
print() print()
print("Calculating per-layer residual directions...") print("Calculating per-layer residual directions...")
print("* Obtaining residual mean for good prompts...") print("* Obtaining residual mean for good prompts...")
good_means = model.get_residuals_mean( good_means = model.get_residuals_mean(
self.good_prompts, good_prompts,
winsorization_quantile=self.settings.winsorization_quantile, winsorization_quantile=self.settings.winsorization_quantile,
) )
print("* Obtaining residual mean for bad prompts...") print("* Obtaining residual mean for bad prompts...")
bad_means = model.get_residuals_mean( bad_means = model.get_residuals_mean(
self.bad_prompts, bad_prompts,
winsorization_quantile=self.settings.winsorization_quantile, winsorization_quantile=self.settings.winsorization_quantile,
) )
+351
View File
@@ -0,0 +1,351 @@
# SPDX-License-Identifier: AGPL-3.0-or-later
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
# Arbitrary-Rank Ablation (ARA) (Weidmann 2026)
# See https://github.com/p-e-w/heretic/pull/211 for more information.
from dataclasses import asdict, dataclass
from typing import Any, cast
import bitsandbytes as bnb
import torch
import torch.linalg as LA
import torch.nn.functional as F
from optuna import Trial
from peft.tuners.lora.layer import Linear
from pydantic import (
BaseModel,
Field,
PositiveInt,
)
from torch import Tensor
from torch.optim import LBFGS
from heretic.config import DatasetSpecification, SingleDatasetSpecification
from heretic.modifier import Context, Modifier, Serializable
from heretic.utils import format_dataset_specification, print
@dataclass
class Parameters(Serializable):
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
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def to_presentation_dict(self) -> dict[str, str]:
return {
name: (f"{value:.4f}" if isinstance(value, float) else f"{value}")
for name, value in asdict(self).items()
}
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "Serializable":
return Parameters(**data)
class Settings(BaseModel):
good_prompts: DatasetSpecification = Field(
default=SingleDatasetSpecification(
dataset="mlabonne/harmless_alpaca",
split="train[:400]",
column="text",
),
description="Dataset of prompts that tend to produce desirable responses.",
)
bad_prompts: DatasetSpecification = Field(
default=SingleDatasetSpecification(
dataset="mlabonne/harmful_behaviors",
split="train[:400]",
column="text",
),
description="Dataset of prompts that tend to produce undesirable responses.",
)
preserve_row_magnitudes: bool = Field(
default=True,
description=(
"Whether to renormalize the rows of the modified matrices to preserve "
"the original matrices' row magnitudes. This is believed to improve "
'intelligence retention (see Lai 2025, "Magnitude-Preserving Orthogonal Ablation").'
),
)
lora_rank: PositiveInt = Field(
default=50,
description=(
"The rank of the LoRA adapter to use. "
'While mathematically, ARA is of "arbitrary" rank, experiments have shown that '
"singular values tend to drop rapidly after a few dozen dimensions, and approximating "
"the full transformation with a LoRA has many practical advantages."
),
)
n_optimization_steps: PositiveInt = Field(
default=5,
description="Number of (outer) L-BFGS optimization steps to perform.",
)
learning_rate: float = Field(
default=1.0,
description="Learning rate to use in the L-BFGS optimizer.",
)
max_iter: PositiveInt = Field(
default=20,
description="Maximum number of (inner) iterations to perform per (outer) L-BFGS optimization step.",
)
history_size: PositiveInt = Field(
default=10,
description="Number of past updates to store for approximating the Hessian matrix in the L-BFGS optimizer.",
)
print_loss: bool = Field(
default=False,
description="Whether to print the loss value for each L-BFGS optimization step.",
)
# 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)
# The objective function at the heart of ARA.
def ara_loss(
good_output: Tensor,
bad_output: Tensor,
new_good_output: Tensor,
new_bad_output: Tensor,
parameters: Parameters,
) -> Tensor:
# 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
)
class ARA(Modifier[Parameters]):
settings: Settings
@property
def reproducible(self) -> bool:
return True
@property
def modifier_name(self) -> str:
return "Arbitrary-Rank Ablation (ARA)"
def init(self, ctx: Context) -> None:
model = ctx.get_model()
print()
print(
f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..."
)
good_prompts = ctx.load_prompts(self.settings.good_prompts)
print(f"* [bold]{len(good_prompts)}[/] prompts loaded")
print()
print(
f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..."
)
bad_prompts = ctx.load_prompts(self.settings.bad_prompts)
print(f"* [bold]{len(bad_prompts)}[/] prompts loaded")
print()
print("Obtaining module I/O for good prompts...")
self.good_module_io = model.get_module_io_batched(good_prompts)
print()
print("Obtaining module I/O for bad prompts...")
self.bad_module_io = model.get_module_io_batched(bad_prompts)
# LoRA B matrices are initialized to zero by default in PEFT,
# so we don't need to do anything manually.
model.apply_lora(self.settings.lora_rank)
def suggest_parameters(self, ctx: Context, trial: Trial) -> Parameters:
layer_count = len(ctx.get_model().get_layers())
start_layer_index = trial.suggest_int(
"start_layer_index",
0,
layer_count // 2,
)
end_layer_index = trial.suggest_int(
"end_layer_index",
layer_count // 2,
layer_count,
)
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.0001,
1.0,
log=True,
)
overcorrect_relative_weight = trial.suggest_float(
"overcorrect_relative_weight",
0.0,
1.3,
)
neighbor_count = trial.suggest_int(
"neighbor_count",
1,
15,
)
return Parameters(
start_layer_index=start_layer_index,
end_layer_index=end_layer_index,
preserve_good_behavior_weight=preserve_good_behavior_weight,
steer_bad_behavior_weight=steer_bad_behavior_weight,
overcorrect_relative_weight=overcorrect_relative_weight,
neighbor_count=neighbor_count,
)
def modify_model(self, ctx: Context, parameters: Parameters) -> None:
model = ctx.get_model()
for layer_index in range(
parameters.start_layer_index,
parameters.end_layer_index,
):
for component, modules in model.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)
# 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:
# Use the original dequantization logic from bitsandbytes.
W_base = cast(
Tensor,
bnb.functional.dequantize_4bit( # ty:ignore[possibly-missing-attribute]
base_weight.data,
quant_state,
).to(torch.float32),
)
# Pre-calculate the original row norms to preserve them.
# See https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration
W_row_norms = cast(
Tensor,
LA.vector_norm(W_base, dim=1, keepdim=True).detach(),
)
# 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)
# Move I/O tensors to the device of the adapter weights.
good_input, good_output = self.good_module_io[layer_index][
component
][module_index]
bad_input, bad_output = self.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)
def objective(A: Tensor, B: Tensor) -> Tensor:
# Calculate effective weight after applying adapter.
W_eff = W_base + (B @ A)
if self.settings.preserve_row_magnitudes:
# Normalize to unit length, then scale by original norms,
# preserving the original row 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
return ara_loss(
good_output,
bad_output,
new_good_output,
new_bad_output,
parameters,
)
optimizer = LBFGS(
[lora_A, lora_B],
lr=self.settings.learning_rate,
max_iter=self.settings.max_iter,
history_size=self.settings.history_size,
line_search_fn="strong_wolfe",
)
def closure() -> Tensor:
optimizer.zero_grad()
loss = objective(lora_A, lora_B)
loss.backward()
return loss
for step in range(self.settings.n_optimization_steps):
loss = optimizer.step(closure)
if self.settings.print_loss:
print(
f"\\[{layer_index}/{component}/{module_index}] Step: {step + 1}, Loss: {loss.item():.6f}"
)
# Free the gradient buffers accumulated during optimization.
# Without this, they persist on the model (one full-size gradient
# per processed weight) and can easily consume tens of GB of VRAM,
# causing out-of-memory errors during the subsequent evaluation.
optimizer.zero_grad(set_to_none=True)
def reset_model(self, ctx: Context) -> None:
model = ctx.get_model()
fast_path = model.reset_model()
if not fast_path:
model.apply_lora(self.settings.lora_rank)
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2 batch_size = 2
max_response_length = 10 max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2 n_trials = 2
n_startup_trials = 1 n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2 batch_size = 2
max_response_length = 10 max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2 n_trials = 2
n_startup_trials = 1 n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2 batch_size = 2
max_response_length = 10 max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2 n_trials = 2
n_startup_trials = 1 n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2 batch_size = 2
max_response_length = 10 max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2 n_trials = 2
n_startup_trials = 1 n_startup_trials = 1
+5
View File
@@ -9,6 +9,11 @@ print_debug_information = true
batch_size = 2 batch_size = 2
max_response_length = 10 max_response_length = 10
modifiers = [
{ plugin = "heretic.modifiers.abliteration.Abliteration" },
]
n_trials = 2 n_trials = 2
n_startup_trials = 1 n_startup_trials = 1