Safety fine-tuning teaches a model to refuse harmful requests as they are usually phrased. Attackers do not phrase them that way. Optimised suffixes, role-play framings, multi-turn escalation and prefilled answers all reach inputs the refusal training never covered, and a model that refuses 99 percent of a benchmark can still fail almost every time against an optimiser. Adversarial training is the model developer's response: put an attacker inside the training loop, find the inputs that currently break the model, and train on exactly those.

The idea comes from image classifiers, where Madry and colleagues framed it in 2017 as a min-max problem. For language models the inner maximisation is much harder, because inputs are discrete tokens and the attack space is effectively unbounded. This article explains the objective, the four families of method used for LLMs, a working embedding-space training step in PyTorch, how to evaluate the result honestly, and the failure modes, chief among them over-refusal and robustness that only holds against the attack you trained on. It assumes the basics of refusal training and of GCG suffix attacks.

The objective and why the inner max is the hard part

Write the model as f with parameters θ, a harmful request x, an attacker perturbation δ from some allowed set Δ, a target harmful completion yh and a refusal yr. Adversarial training solves:

min over theta of  E over (x, y_r, y_h) [  max over delta in Delta  L_defence(f_theta(x + delta), y_r, y_h)  ]
                 + gamma * E over benign (x_b, y_b) [ L_sft(f_theta(x_b), y_b) ]

The inner max finds the worst perturbation for the current model; the outer min fixes it. The benign term is not optional. Without it the cheapest way to minimise the defence loss is to refuse everything, and that is what an unregularised run converges to.

Everything practical hinges on Δ, the attacker's allowed moves. Make it too narrow and the model becomes robust to a toy threat. Make it require a full discrete search such as GCG at every step and a single training run costs orders of magnitude more than ordinary fine-tuning. The four families below are four different answers to that trade-off.

Adversarial training: an attacker inside the training loopharmful behavioursprompt + target + refusalbenign datautility + borderline setinner attackPGD on embeddings or latentsadversarial inputdefence lossestoward refusal, away from targetutility lossupdate modelLoRA or fullnew weightsheld-out evaluationunseen attacks, over-refusal, utilityThe inner attack must be cheap enough to run every step, and the evaluation must use attacks the loop never saw.
The training loop. Harmful data feeds the inner attack, benign data keeps the model useful, and evaluation sits outside the loop on attacks it never saw.

Four families of method

1. Data-level adversarial training. Human red teams and automated attack generators produce prompts that succeed against the current model; these are paired with refusals and folded into the next round of supervised or preference training. This is how most production safety pipelines started. It is cheap per step and covers realistic phrasing, but the attacks are found against an older model, so each round lags the one before. Automated red teaming covers the generator side.

2. Discrete attacks in the loop: R2D2. The HarmBench paper (Mazeika et al., 2024) introduced R2D2, Robust Refusal Dynamic Defense. It keeps a persistent pool of adversarial test cases and advances each by only a few GCG steps per training iteration, amortising the search across training instead of restarting it. The loss pushes the model away from the harmful target string, toward a refusal, and keeps an ordinary instruction-tuning term for utility. It showed that training against GCG works against GCG, at a substantial compute cost.

3. Continuous embedding attacks. Xhonneux et al. (2024) replaced the discrete search with projected gradient steps on the input embeddings of an adversarial suffix. Embeddings are continuous, so a handful of gradient steps finds a strong perturbation, making each training step a small constant factor more expensive than plain fine-tuning. The perturbation is not constrained to real tokens, which makes it a stronger attacker than any token sequence; the bet is that robustness to the stronger attacker transfers to discrete ones.

4. Latent adversarial training. Casper et al. (2024) and Sheshadri et al. (2024) move the perturbation deeper, into hidden activations at a chosen layer. Attacking the residual stream directly targets the internal features that harmful behaviour uses, rather than any one way of reaching them from the input, and the targeted variant aims to remove persistent harmful behaviours that ordinary refusal training only suppresses.

A related but distinct method is circuit breakers (Zou et al., 2024), which trains the model to reroute internal representations of harmful content into an unusable state instead of training against an attacker. It is often compared with adversarial training, and the two can be combined.

A training step in PyTorch

The step below implements the continuous-embedding approach. The adversarial suffix occupies fixed positions between the prompt and the completion, so its perturbation can be reused when scoring both the harmful target and the refusal. The harmful targets are short affirmative prefixes in the HarmBench style, such as a placeholder beginning "Sure, here is", not harmful content: the attack only has to make the model start to comply. For clarity the batch is assumed unpadded and equal-length.

import torch
import torch.nn.functional as F

def completion_nll(model, embed, pre, adv, comp_ids):
    """Mean NLL of comp_ids given prompt embeddings `pre` followed by suffix embeddings `adv`."""
    x = torch.cat([pre, adv, embed(comp_ids)], dim=1)
    logits = model(inputs_embeds=x).logits
    start = pre.shape[1] + adv.shape[1]
    pred = logits[:, start - 1:-1].float()            # position i predicts token i + 1
    return F.cross_entropy(pred.reshape(-1, pred.shape[-1]), comp_ids.reshape(-1))

def embedding_pgd(model, embed, pre, adv0, target_ids, eps, alpha, steps):
    """Find a per-token L2-bounded perturbation that makes the harmful target likely."""
    delta = torch.zeros_like(adv0, requires_grad=True)
    for _ in range(steps):
        loss = completion_nll(model, embed, pre, adv0 + delta, target_ids)
        (g,) = torch.autograd.grad(loss, delta)       # gradient w.r.t. delta only
        with torch.no_grad():
            delta -= alpha * g / (g.norm(dim=-1, keepdim=True) + 1e-12)
            scale = (eps / (delta.norm(dim=-1, keepdim=True) + 1e-12)).clamp(max=1.0)
            delta *= scale                           # project back into the eps ball
    return (adv0 + delta).detach()

def train_step(model, embed, opt, harm, benign_loss_fn, cfg):
    with torch.no_grad():
        pre = embed(harm.prompt_ids)
        adv0 = embed(harm.suffix_init_ids)            # e.g. twenty copies of "!"
    adv = embedding_pgd(model, embed, pre, adv0, harm.target_ids,
                        cfg.eps, cfg.alpha, cfg.attack_steps)
    toward = completion_nll(model, embed, pre, adv, harm.refusal_ids)
    away = -completion_nll(model, embed, pre, adv, harm.target_ids).clamp(max=cfg.away_cap)
    utility = benign_loss_fn(model)                   # ordinary SFT loss on benign data
    loss = toward + cfg.beta * away + cfg.gamma * utility
    opt.zero_grad(set_to_none=True)
    loss.backward()
    opt.step()
    return dict(toward=toward.item(), away=away.item(), utility=utility.item())

Three details carry the method. The attack differentiates only with respect to δ, so it never touches weights. The away term is clamped: an unbounded negative log-likelihood can be driven to infinity by degrading the whole language model, and the clamp stops that once the harmful target is improbable enough. And ε is per token and relative to the embedding scale; set it as a fraction of the average embedding norm, measured once, rather than as an absolute number borrowed from a paper on a different model.

Trace one step to see the cost. Take a batch of 8 harmful prompts of 40 tokens, a 20-token suffix and 12-token targets and refusals, on a model with 4,096-dimensional embeddings. The perturbation δ has shape 8 x 20 x 4,096. Ten PGD steps cost ten forward and ten backward passes over 72-token sequences; the defence step adds two more forwards and one combined backward, plus the benign batch. Compared with plain fine-tuning on the same data, a step is therefore several times more expensive, which is affordable, where a full GCG search per example would be thousands of forward passes. Log the three loss terms separately: a healthy run shows the toward loss falling, the away term saturating at its cap, and the utility loss staying flat. A rising utility loss is the earliest warning that the model is buying robustness with capability.

Gradient checkpointing and LoRA keep memory manageable, because the attack's backward passes need activations but no optimiser state. Run the attack in the same precision as training; an attack computed in a lower precision than the defence is weaker than it looks.

Data: four sets and their balance

Four datasets go into a run, and their balance decides the result more than the attack does.

  • Harmful behaviours, each with a refusal and a short affirmative target. Cover your policy's categories evenly; the model learns the distribution you give it.
  • Benign instructions for the utility loss, ideally drawn from the same distribution as your normal fine-tuning data.
  • Borderline benign requests that look risky but are fine, such as questions about the history of a weapon or how a poison is treated medically, with helpful answers. These anchor the decision boundary and are the main defence against over-refusal.
  • Held-out attacks kept entirely outside training: other optimisers, multi-turn scripts, prefilling, paraphrase and translation attacks. These are your test set.

A worked configuration for a first run on a 7B to 8B chat model: LoRA on attention and MLP projections, a 20-token suffix, ten PGD steps per batch, equal numbers of harmful and benign examples per batch, and a borderline set of a few hundred examples mixed into the benign stream. Treat every value as a starting point to sweep, and change one at a time.

Evaluating without fooling yourself

The central risk of adversarial training is that it appears to work. A model trained against embedding PGD will resist embedding PGD; the question is whether it resists anything else. Evaluate along three axes and report all three together:

  1. Robustness to unseen, adaptive attacks. Run discrete attacks such as GCG, multi-turn escalation and prefilling against the final model, with a budget at least as large as the one you trained against. An attack tuned against the base model and replayed is not adaptive. See universal adversarial suffixes for how to measure transfer honestly.
  2. Over-refusal. Measure the refusal rate on benign-but-scary prompts with a public set such as XSTest plus your own borderline data. A robustness gain bought with a large jump here is usually not a gain.
  3. Utility. Run your standard capability suite and a sample of real traffic. Small drops are expected; large ones mean the utility weight γ is too low.

Also test robustness after fine-tuning. Safety behaviour of any kind can be removed by further training on a small dataset, as covered in fine-tuning erodes safety training, so state the threat model: adversarial training raises the cost of input attacks, it does not protect open weights.

Failure modes

  • Robust overfitting. Robustness to the training attack keeps rising while held-out robustness peaks and falls. Checkpoint frequently and select on held-out attacks.
  • Gradient masking. The model learns to make the attacker's gradients uninformative rather than becoming robust. The symptom is that gradient attacks fail while gradient-free or transfer attacks succeed.
  • Refuse-everything collapse. The defence loss is satisfied by refusing all inputs. Watch the borderline refusal rate every evaluation, not just at the end.
  • Language-model damage from the away loss. Unclamped, it raises perplexity on everything. Track benign perplexity as a training curve.
  • Evaluation contamination. Test attacks that leak into training data inflate every number. Keep held-out attacks in a separate store with access control.
  • Mismatched ε. Too small and the attacker is toothless; too large and the perturbation leaves the region of real inputs, teaching nothing useful.

Trade-offs

MethodCost per stepStrengthMain weakness
Data-level red-team roundsSame as SFTRealistic phrasingLags the current model
R2D2 with GCG in the loopHighDirectly targets token attacksCompute; narrow to GCG-like threats
Continuous embedding PGDA few extra passesCheap and strong attackerTransfer to discrete attacks must be shown
Latent adversarial trainingA few extra passesTargets internal featuresLayer and ε choice are delicate
Circuit breakers (related)ModerateActs on harmful representationsDifferent objective; still needs adaptive testing

None of these replaces the deployment-side controls in a layered jailbreak defence: input and output classifiers, rate limits and tool authorisation still matter when the model is wrong.

What to do next

  1. Write down the threat model: which attacks, what budget, and whether weights are public.
  2. Build the four datasets, and lock the held-out attack set away before any training run.
  3. Measure your base model on all three axes, held-out attacks, over-refusal and utility, so every later number has a baseline.
  4. Implement the embedding-PGD step above with LoRA, calibrate ε against the embedding norm, and sweep β and γ.
  5. Select checkpoints on held-out attacks and borderline refusal together, never on training-attack loss.
  6. Re-run adaptive attacks against the final model and record the results next to the base model's in your model card.
Key takeaway: Adversarial training puts an attacker inside the safety fine-tuning loop so the model learns to refuse the inputs that currently break it. Continuous embedding and latent attacks make it affordable, a benign and borderline utility term keeps the model useful, and the result only counts when it holds against adaptive attacks the loop never saw, without a jump in over-refusal.