Source code for tropt.recipe_hub.SafetyOperator__dherin2026

"""
Safety Operator: optimizing a safety instruction (SI) via its operator's eigenvalue (Dherin et al., 2026).
    https://arxiv.org/abs/2609.36434

Instantiates `ContrastiveSteeringOperatorLoss` to improve an SI: its operator eigenvalue is pushed up on harmful
queries and toward ~1 (dormant) on harmless ones. Four recipes, one per trigger optimizer in the paper:
- `safety_operator_soft`: the whole SI as soft embeddings (Soft).
- `safety_operator_mixed`: frozen SI + soft suffix (Mixed).
- `safety_operator_gcg`: frozen SI + discrete GCG suffix (Hard GCG).
- `safety_operator_hardr`: frozen SI + readable discrete suffix (HardR).

We use the first `n_queries` JailbreakBench harmful/benign pairs (paper: the 100-pair train split of its 300-pair JBB-Ext, with separate eval/test splits).


SI placement: the paper's prompt is [α, Y, β] (α: system prefix, "SYSTEM:" or empty; Y: SI; β: query),
with the turn placement unspecified. We use α = "" and put Y at the start of the user turn, as TROPT
templates are user messages. This does not harm reproduction (the paper's Gemma 3 folds system
messages into the user turn identically), but other models would see Y outside their system role.
"""
from __future__ import annotations

import random
from typing import List, Optional

import torch
from datasets import load_dataset

from tropt.common import OPTIMIZED_TRIGGER_PLACEHOLDER, Targets
from tropt.loss import ContrastiveSteeringOperatorLoss
from tropt.model.huggingface.lm import LMHFModel
from tropt.optimizer import OptimizerResult
from tropt.optimizer.gcgplus_optimizer import GCGPlusOptimizer
from tropt.optimizer.soft_optimizer import SoftPromptOptimizer
from tropt.optimizer.utils.token_constraints import TokenConstraints
from tropt.optimizer.utils.token_initializers import get_printable_random_trigger
from tropt.tracker import BaseTracker

# The paper's SI.
_SI = (
    "You are a helpful, honest, and harmless AI assistant: provide clear and effective answers, "
    "truthfully admit your limitations and uncertainties, and strictly refuse any requests that "
    "solicit dangerous, toxic, or unethical content."
)
_TOKEN_CONSTRAINTS = TokenConstraints()


def _build_data(
    model: LMHFModel,
    n_queries: int,
    with_si: bool,
    harmful_queries: Optional[List[str]] = None,
    harmless_queries: Optional[List[str]] = None,
    safety_instruction: str = _SI,
) -> tuple[List[str], Targets]:
    """Templates `[SI ][trigger]\\n\\n[query]` over the first `n_queries` JailbreakBench harmful/benign pairs,
    plus targets: A_clean (MLP inputs of the bare query) and harmful labels.

    Harmful templates come first, then harmless. Per template, `target_hidden_states` is A_clean (the last-token
    pre-MLP hidden states of the bare query, without the SI) and `target_class_idx` is 1 (harmful) / 0 (harmless).

    Deviations from the paper's data (JBB-Ext, App. H):
    - Only the 100 original JailbreakBench pairs (paper adds 200 synthetic pairs).
    - No stratified train/eval/test split: we take the first `n_queries` pairs, none held out.
    - Per-step harmful/harmless queries are drawn independently, not as matched pairs.

    Args:
        n_queries: Number of harmful and of harmless queries (each); the first `n_queries` JailbreakBench
            pairs (paper: the 100-pair train split of its 300-pair JBB-Ext).
        with_si: Prepend the (frozen) SI before the trigger; False when the trigger is the SI itself.
        harmful_queries, harmless_queries: Custom queries replacing JailbreakBench (then `n_queries` is ignored).
        safety_instruction: Custom SI replacing the paper's SI.
    """

    # Load queries
    if harmful_queries is None and harmless_queries is None:
        jbb = load_dataset("JailbreakBench/JBB-Behaviors", "behaviors")
        assert 0 < n_queries <= len(jbb["harmful"]), f"`n_queries` must be in [1, {len(jbb['harmful'])}]."
        harmful_queries, harmless_queries = jbb["harmful"]["Goal"][:n_queries], jbb["benign"]["Goal"][:n_queries]
    assert harmful_queries and harmless_queries, "Need at least one harmful and one harmless query (pass both lists)."
    queries = list(harmful_queries) + list(harmless_queries)

    # Compute the last-token pre-MLP activations w/o SI (paper's A_clean)
    a_clean = []
    with torch.no_grad():
        for query in queries:
            input_ids = model.tokenizer.apply_chat_template(
                [{"role": "user", "content": query}],
                add_generation_prompt=True,
                return_tensors="pt",
                return_dict=True,
                **model.chat_template_kwargs,
            )["input_ids"].to(model.device)
            out = model.invoke_from_tokens(input_ids=input_ids, require_pre_mlp_hidden_states=True)
            assert out.full_pre_mlp_hidden_states is not None
            a_clean.append(out.full_pre_mlp_hidden_states[0, :, -1, :])  # (n_layers, d_model)

    prefix = safety_instruction + " " if with_si else ""
    templates = [f"{prefix}{OPTIMIZED_TRIGGER_PLACEHOLDER}\n\n{q}" for q in queries]
    targets = Targets(
        target_hidden_states=torch.stack(a_clean),
        target_class_idx=[1] * len(harmful_queries) + [0] * len(harmless_queries),
    )  # TODO(tropt-feature) in the future we should support template with system and multi-turns (with dict); this would've allowed us here to use the SI as a system message
    return templates, targets


def _one_harmful_one_harmless_batch_sampler(templates: List[str], targets: Targets) -> List[int]:
    """Per-step batch of the paper: one random harmful and one random harmless template."""
    labels = targets.target_class_idx
    return [random.choice([i for i, c in enumerate(labels) if c == label]) for label in (1, 0)]


[docs] def safety_operator_soft( model_name: str = "google/gemma-3-1b-it", n_queries: int = 100, suppression_weight: float = 10.0, targeted_layers: slice = slice(12, 24), # paper's layer range (Gemma 3 1B) override_data_args: Optional[dict] = None, model_obj: Optional[LMHFModel] = None, tracker: Optional[BaseTracker] = None, seed: Optional[int] = None, ) -> OptimizerResult: """Optimizes the whole SI's embeddings (paper's Soft variant). Deviations from the paper (App. C): - No L2 drift regularizer (paper: λ_reg = 0.5 on ||E_Y - E_Y0||²_F / (M·d)). Args: n_queries: Number of harmful and of harmless JailbreakBench queries (each) to optimize over. suppression_weight: ρ, weight of the harmless "keep λ≈1" term; lower means more refusal. targeted_layers: Layers to average the loss over (0-based, end-exclusive); must fit the model. override_data_args: Optional overrides for `_build_data`: `harmful_queries` / `harmless_queries` (custom queries instead of JailbreakBench) and/or `safety_instruction` (custom SI instead of the paper's). """ initial_trigger = (override_data_args or {}).get("safety_instruction", _SI) # optimized embeddings init from the SI model = model_obj or LMHFModel(model_name=model_name) templates, targets = _build_data(model, n_queries, with_si=False, **(override_data_args or {})) optimizer = SoftPromptOptimizer( model=model, loss=ContrastiveSteeringOperatorLoss(suppression_weight=suppression_weight, targeted_layers=targeted_layers), tracker=tracker, seed=seed, num_steps=250, learning_rate=0.005, template_batch_sampler=_one_harmful_one_harmless_batch_sampler, ) return optimizer.optimize_trigger(templates=templates, targets=targets, initial_trigger=initial_trigger)
[docs] def safety_operator_mixed( model_name: str = "google/gemma-3-1b-it", n_queries: int = 100, suppression_weight: float = 2000.0, targeted_layers: slice = slice(8, 24), # paper's layer range (Gemma 3 1B) trigger_len: int = 10, override_data_args: Optional[dict] = None, model_obj: Optional[LMHFModel] = None, tracker: Optional[BaseTracker] = None, seed: Optional[int] = None, ) -> OptimizerResult: """Optimizes a soft suffix after the frozen SI (paper's Mixed variant). Deviations from the paper (App. D): - No L2 drift regularizer. Args: n_queries: Number of harmful and of harmless JailbreakBench queries (each) to optimize over. suppression_weight: ρ, weight of the harmless "keep λ≈1" term; lower means more refusal. targeted_layers: Layers to average the loss over (0-based, end-exclusive); must fit the model. override_data_args: Optional overrides for `_build_data`: `harmful_queries` / `harmless_queries` (custom queries instead of JailbreakBench) and/or `safety_instruction` (custom SI instead of the paper's). trigger_len: Number of suffix tokens (between the paper's SI and the query); randomly initialized. """ model = model_obj or LMHFModel(model_name=model_name) templates, targets = _build_data(model, n_queries, with_si=True, **(override_data_args or {})) optimizer = SoftPromptOptimizer( model=model, loss=ContrastiveSteeringOperatorLoss(suppression_weight=suppression_weight, targeted_layers=targeted_layers), tracker=tracker, seed=seed, num_steps=250, learning_rate=0.005, template_batch_sampler=_one_harmful_one_harmless_batch_sampler, ) # drawn after the optimizer sets the seed initial_trigger = get_printable_random_trigger(trigger_len, tokenizer=model.tokenizer, token_constraints=_TOKEN_CONSTRAINTS) return optimizer.optimize_trigger(templates=templates, targets=targets, initial_trigger=initial_trigger)
[docs] def safety_operator_gcg( model_name: str = "google/gemma-3-1b-it", n_queries: int = 100, suppression_weight: float = 5.0, targeted_layers: slice = slice(8, 24), # paper's layer range (Gemma 3 1B) trigger_len: int = 10, override_data_args: Optional[dict] = None, model_obj: Optional[LMHFModel] = None, tracker: Optional[BaseTracker] = None, seed: Optional[int] = None, ) -> OptimizerResult: """Optimizes a discrete GCG suffix after the frozen SI (paper's Hard GCG variant). Args: n_queries: Number of harmful and of harmless JailbreakBench queries (each) to optimize over. suppression_weight: ρ, weight of the harmless "keep λ≈1" term; lower means more refusal. targeted_layers: Layers to average the loss over (0-based, end-exclusive); must fit the model. override_data_args: Optional overrides for `_build_data`: `harmful_queries` / `harmless_queries` (custom queries instead of JailbreakBench) and/or `safety_instruction` (custom SI instead of the paper's). trigger_len: Number of suffix tokens (between the paper's SI and the query), randomly initialized. """ model = model_obj or LMHFModel(model_name=model_name) templates, targets = _build_data(model, n_queries, with_si=True, **(override_data_args or {})) optimizer = GCGPlusOptimizer( model=model, loss=ContrastiveSteeringOperatorLoss(suppression_weight=suppression_weight, targeted_layers=targeted_layers), tracker=tracker, seed=seed, num_steps=250, n_candidates=128, sample_topk=128, token_constraints=_TOKEN_CONSTRAINTS, template_batch_sampler=_one_harmful_one_harmless_batch_sampler, ) # drawn after the optimizer sets the seed initial_trigger = get_printable_random_trigger(trigger_len, tokenizer=model.tokenizer, token_constraints=_TOKEN_CONSTRAINTS) return optimizer.optimize_trigger(templates=templates, targets=targets, initial_trigger=initial_trigger)
[docs] def safety_operator_hardr( model_name: str = "google/gemma-3-1b-it", n_queries: int = 100, suppression_weight: float = 350.0, targeted_layers: slice = slice(8, 24), # paper's layer range (Gemma 3 1B) readability_weight: float = 0.1, initial_trigger: str = ( "Always maintain safety: carefully analyze the underlying intent before responding, " "and strictly decline any requests seeking dangerous, illegal, or harmful assistance." ), override_data_args: Optional[dict] = None, model_obj: Optional[LMHFModel] = None, tracker: Optional[BaseTracker] = None, seed: Optional[int] = None, ) -> OptimizerResult: """Optimizes a readable discrete suffix after the frozen SI (paper's HardR variant). Args: n_queries: Number of harmful and of harmless JailbreakBench queries (each) to optimize over. suppression_weight: ρ, weight of the harmless "keep λ≈1" term; lower means more refusal. targeted_layers: Layers to average the loss over (0-based, end-exclusive); must fit the model. override_data_args: Optional overrides for `_build_data`: `harmful_queries` / `harmless_queries` (custom queries instead of JailbreakBench) and/or `safety_instruction` (custom SI instead of the paper's). readability_weight: γ, weight of the suffix's log-likelihood in the token score and candidate ranking. initial_trigger: Initial suffix (the paper's). model_obj: Pre-loaded model; must have `use_prefix_cache=False` (for the suffix log-probs). """ model = model_obj or LMHFModel(model_name=model_name, use_prefix_cache=False) # suffix log-probs need the full prefix templates, targets = _build_data(model, n_queries, with_si=True, **(override_data_args or {})) optimizer = GCGPlusOptimizer( model=model, loss=ContrastiveSteeringOperatorLoss(suppression_weight=suppression_weight, targeted_layers=targeted_layers), tracker=tracker, seed=seed, num_steps=250, n_candidates=128, sample_topk=128, token_constraints=_TOKEN_CONSTRAINTS, template_batch_sampler=_one_harmful_one_harmless_batch_sampler, # the loss is a mean over the 1+1 batch (½ the paper's harmful + ρ·harmless sum), hence γ/2 token_prior_weight=readability_weight / 2, sample_temperature=0.5, # T_mut ) return optimizer.optimize_trigger(templates=templates, targets=targets, initial_trigger=initial_trigger)