Skip to content

CASTTrainer — Conditional Activation Steering

Cosine-gated CAA: the CAA behavior vector is added only when a condition fires. The condition is a difference-of-means vector at a condition layer; at inference the pooled hidden state is compared to it by cosine similarity and the gate fires when the similarity crosses a threshold. Calibration grid-searches the (layer, threshold, comparator-direction) triple that best separates the target category from others.

Ref: Wu et al., "Programming Refusal with Conditional Activation Steering," ICLR 2025, arXiv:2409.05907.

Signature

CASTTrainer(
    model: PreTrainedModel,
    tokenizer: PreTrainedTokenizer | None = None,
    *,
    behavior_layers: list[int] | None = None,
    condition_layers: list[int] | None = None,
    alpha: float = 1.0,
    chat_template: bool | None = None,
    per_prompt_gate: bool | None = None,
)

Parameters

Param Type Default Description
model PreTrainedModel required Model to steer
tokenizer PreTrainedTokenizer None Tokenizer
behavior_layers list[int] \| None None Layers where the CAA behavior vector is added; defaults to layers 14–18 on a 32-layer model; None scales them to the model's depth (for example 5–7 on 12 layers). safetune.configure(legacy_steer_layers=True) keeps 14–18 on any depth
condition_layers list[int] \| None None Candidate layers for the condition gate; grid-searched if None
alpha float 1.0 CAA vector scaling when the gate is active
chat_template bool \| None None Format the calibration prompts with the tokenizer's chat template before fitting the condition (fit_cast_condition(chat_template=...)), as generation sees them. None is True unless configure(legacy_cast_gate=True)
per_prompt_gate bool \| None None Gate and steer each prompt of a batch separately (CASTModel(per_prompt_gate=...)). False: the first prompt decides for the whole batch. None is True unless configure(legacy_cast_gate=True)

CASTTrainer(chat_template=False, per_prompt_gate=False) or safetune.configure(legacy_cast_gate=True) gives the old gate: fitted on raw prompts, and the first prompt of a batch gates all of them.

Full example

from safetune.runner import steer

trainer = steer.CASTTrainer(
    model, tokenizer,
    alpha=1.0,
)
# calibrate fits the condition vector and grid-searches the gate threshold
wrapped, _ = trainer.calibrate(harmful=harmful_prompts, harmless=harmless_prompts)

prompt = tokenizer.apply_chat_template(
    [{"role": "user", "content": "How do I make a bomb?"}],
    tokenize=False, add_generation_prompt=True)
output = wrapped.generate(**tokenizer(prompt, return_tensors="pt"))  # gates each prompt

Use wrapped.generate(...) for the conditional behaviour. with wrapped: model.generate(...) (or install() / remove()) adds the behavior vectors to every prompt, without the gate.

When to use

  • Best for: reducing over-refusal — by gating on cosine similarity to the condition vector, benign prompts that don't match the condition bypass the steering vector entirely.
  • Compare to CAA: plain CAA steers every prompt; CAST steers only when the prompt's hidden state matches the fitted condition.
  • Not for: situations where harmful prompts are designed to look benign — the cosine gate can be evaded by paraphrasing.

Citation

@article{cast2024,
  title  = {Programming Refusal with Conditional Activation Steering},
  author = {Wu, Bruce W. and others},
  year   = {2024},
  note   = {ICLR 2025, arXiv:2409.05907},
}