Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
231 changes: 223 additions & 8 deletions pyrit/executor/promptgen/gcg/attack/base/attack_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,13 @@

from __future__ import annotations

import inspect
import json
import logging
import math
import random
import time
from copy import deepcopy
from copy import copy, deepcopy
from dataclasses import dataclass
from enum import Enum
from typing import TYPE_CHECKING, Any, cast
Expand Down Expand Up @@ -41,6 +42,8 @@

from transformers import PreTrainedModel, PreTrainedTokenizerBase

from pyrit.executor.promptgen.gcg.extension_protocols import LossFunction

logger = logging.getLogger(__name__)

_DEFAULT_TEST_PREFIXES: list[str] = [
Expand Down Expand Up @@ -559,12 +562,17 @@ def grad(self, model: Any) -> torch.Tensor:
raise NotImplementedError("Gradient function not yet implemented")

@torch.no_grad() # type: ignore[misc, untyped-decorator, unused-ignore]
def logits(self, model: Any, test_controls: Any = None, return_ids: bool = False) -> Any:
def _build_candidate_batch(
self,
model: Any,
test_controls: Any = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""
Compute logits for one or more candidate controls.
Build candidate token ids and their optional padding mask.

Returns:
Any: Model logits, optionally paired with their token ids.
tuple[torch.Tensor, torch.Tensor | None]: Candidate token ids and
an attention mask when string controls require padding.

Raises:
ValueError: If candidate controls have an invalid type or shape.
Expand Down Expand Up @@ -610,15 +618,221 @@ def logits(self, model: Any, test_controls: Any = None, return_ids: bool = False
self.input_ids.unsqueeze(0).repeat(test_ids.shape[0], 1).to(model.device), 1, locs, test_ids
)
attn_mask = (ids != pad_tok).type(ids.dtype) if pad_tok >= 0 else None
return ids, attn_mask

@torch.no_grad() # type: ignore[misc, untyped-decorator, unused-ignore]
def logits(
self,
model: Any,
test_controls: Any = None,
return_ids: bool = False,
logits_to_keep: torch.Tensor | None = None,
) -> Any:
"""
Compute logits for one or more candidate controls.

Returns:
Any: Model logits, optionally paired with their token ids.

Raises:
ValueError: If candidate controls have an invalid type or shape.
"""
ids, attn_mask = self._build_candidate_batch(model, test_controls)

model_kwargs: dict[str, Any] = {"input_ids": ids, "attention_mask": attn_mask}
if logits_to_keep is not None:
model_kwargs["logits_to_keep"] = logits_to_keep

if return_ids:
del locs, test_ids
return model(input_ids=ids, attention_mask=attn_mask).logits, ids
del locs, test_ids
logits = model(input_ids=ids, attention_mask=attn_mask).logits
return model(**model_kwargs).logits, ids
logits = model(**model_kwargs).logits
del ids
return logits

@staticmethod
def _expand_prefix_cache(prefix_cache: Any, batch_size: int) -> Any:
"""
Return an independently mutable, batch-expanded view of a model KV cache.

Returns:
Any: A cache whose batch dimension is expanded to ``batch_size``.

Raises:
ValueError: If the source cache contains more than one sequence.
TypeError: If the model returned an unsupported cache structure.
"""

def contains_non_scalar_tensor(value: Any) -> bool:
if isinstance(value, torch.Tensor):
return value.ndim > 0
if isinstance(value, dict):
return any(contains_non_scalar_tensor(item) for item in value.values())
if isinstance(value, (tuple, list)):
return any(contains_non_scalar_tensor(item) for item in value)
return False

if hasattr(prefix_cache, "layers"):
expanded_cache = copy(prefix_cache)
expanded_layers = []
for layer in prefix_cache.layers:
keys = getattr(layer, "keys", None)
values = getattr(layer, "values", None)
if not isinstance(keys, torch.Tensor) or not isinstance(values, torch.Tensor):
raise TypeError(f"Unsupported prefix-cache layer type: {type(layer)!r}")
try:
additional_state = (value for name, value in vars(layer).items() if name not in {"keys", "values"})
except TypeError as exc:
raise TypeError(f"Unsupported prefix-cache layer type: {type(layer)!r}") from exc
if any(contains_non_scalar_tensor(value) for value in additional_state):
raise TypeError(f"Unsupported state in prefix-cache layer type: {type(layer)!r}")
if keys.ndim == 0 or values.ndim == 0 or keys.shape[0] != 1 or values.shape[0] != 1:
raise ValueError("Prefix cache must be computed for exactly one sequence")

expanded_layer = copy(layer)
expanded_layer.keys = keys.expand(batch_size, *keys.shape[1:])
expanded_layer.values = values.expand(batch_size, *values.shape[1:])
expanded_layers.append(expanded_layer)
expanded_cache.layers = expanded_layers
return expanded_cache

if isinstance(prefix_cache, (tuple, list)):
expanded_legacy_cache = []
for layer in prefix_cache:
if not isinstance(layer, (tuple, list)) or not all(isinstance(value, torch.Tensor) for value in layer):
raise TypeError(f"Unsupported prefix-cache layer type: {type(layer)!r}")
if any(value.ndim == 0 or value.shape[0] != 1 for value in layer):
raise ValueError("Prefix cache must be computed for exactly one sequence")
expanded_legacy_cache.append(tuple(value.expand(batch_size, *value.shape[1:]) for value in layer))
return tuple(expanded_legacy_cache)

raise TypeError(f"Unsupported prefix-cache type: {type(prefix_cache)!r}")

def _loss_with_prefix_cache(
self,
model: Any,
test_controls: Any,
loss_function: Any,
logit_positions: torch.Tensor,
) -> torch.Tensor | None:
"""
Score candidates after evaluating their invariant prefix once.

Returns:
torch.Tensor | None: One scalar loss per candidate control, or ``None``
when the model's cache cannot be safely batch-expanded.
"""
token_ids, attention_mask = self._build_candidate_batch(model, test_controls)
prefix_length = self._control_slice.start - 1
prefix_kwargs: dict[str, Any] = {
"input_ids": token_ids[:1, :prefix_length],
"use_cache": True,
"return_dict": True,
"logits_to_keep": 1,
}
if attention_mask is not None:
prefix_kwargs["attention_mask"] = attention_mask[:1, :prefix_length]
prefix_output = model(**prefix_kwargs)
Comment thread
romanlutz marked this conversation as resolved.
try:
prefix_cache = self._expand_prefix_cache(
getattr(prefix_output, "past_key_values", None), token_ids.shape[0]
)
except (TypeError, ValueError):
del prefix_output, token_ids
return None
del prefix_output

suffix_kwargs: dict[str, Any] = {
"input_ids": token_ids[:, prefix_length:],
"past_key_values": prefix_cache,
"use_cache": False,
"return_dict": True,
"logits_to_keep": logit_positions - prefix_length,
}
if attention_mask is not None:
suffix_kwargs["attention_mask"] = attention_mask
logits = model(**suffix_kwargs).logits
del prefix_cache

try:
result: torch.Tensor = loss_function.compute_loss_from_selected_logits(
logits=logits,
token_ids=token_ids,
target_slice=self._target_slice,
control_slice=self._control_slice,
)
return result
finally:
del logits, token_ids

def loss(
self,
model: Any,
test_controls: Any,
loss_function: LossFunction,
*,
use_prefix_cache: bool = False,
) -> torch.Tensor:
"""
Compute per-candidate loss without returning full logits from the worker.

The model forward pass and loss calculation stay in the process that owns
the model. Only the batch-sized loss tensor crosses the worker boundary.

Returns:
torch.Tensor: One scalar loss per candidate control.
"""
selective_loss = cast("Any", loss_function)
try:
forward_parameters = inspect.signature(model.forward).parameters
except (TypeError, ValueError):
forward_parameters = {}

supports_selective_logits = "logits_to_keep" in forward_parameters
supports_prefix_cache = "past_key_values" in forward_parameters and "use_cache" in forward_parameters

if supports_selective_logits and hasattr(selective_loss, "get_required_logit_positions"):
logit_positions = selective_loss.get_required_logit_positions(
target_slice=self._target_slice,
control_slice=self._control_slice,
device=model.device,
)
if use_prefix_cache and supports_prefix_cache and self._control_slice.start > 1:
cached_loss = self._loss_with_prefix_cache(
model,
test_controls,
selective_loss,
logit_positions,
)
if cached_loss is not None:
return cached_loss
logits, token_ids = self.logits(
model,
test_controls,
return_ids=True,
logits_to_keep=logit_positions,
)
try:
result: torch.Tensor = selective_loss.compute_loss_from_selected_logits(
logits=logits,
token_ids=token_ids,
target_slice=self._target_slice,
control_slice=self._control_slice,
)
return result
finally:
del logits, token_ids

logits, token_ids = self.logits(model, test_controls, return_ids=True)
try:
return loss_function.compute_loss(
logits=logits,
token_ids=token_ids,
target_slice=self._target_slice,
control_slice=self._control_slice,
)
finally:
del logits, token_ids

def target_loss(self, logits: torch.Tensor, ids: torch.Tensor) -> torch.Tensor:
"""
Compute unreduced cross-entropy loss over target tokens.
Expand Down Expand Up @@ -2051,6 +2265,7 @@ class ModelWorkerOperation(str, Enum):

GRAD = "grad"
LOGITS = "logits"
LOSS = "loss"
CONTRAST_LOGITS = "contrast_logits"
TEST = "test"
TEST_LOSS = "test_loss"
Expand Down
57 changes: 39 additions & 18 deletions pyrit/executor/promptgen/gcg/attack/gcg/candidate_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
ModelWorkerOperation,
PromptManager,
)
from pyrit.executor.promptgen.gcg.default_implementations import CrossEntropyLoss
from pyrit.executor.promptgen.gcg.extension_protocols import LossFunction


Expand Down Expand Up @@ -46,6 +47,7 @@ def __init__(
prompts: list[PromptManager],
loss_function: LossFunction,
main_device: torch.device,
use_prefix_cache: bool = False,
) -> None:
"""
Initialize the candidate evaluator.
Expand All @@ -55,6 +57,8 @@ def __init__(
prompts: List of prompt managers associated with each worker.
loss_function: Loss function protocol used to compute candidate losses.
main_device: PyTorch device on which aggregate losses are stored and accumulated.
use_prefix_cache: Whether built-in loss evaluation may reuse the
invariant prefix KV cache.

Raises:
ValueError: If workers list is empty, or if worker and prompt manager counts mismatch,
Expand All @@ -73,6 +77,8 @@ def __init__(
self._prompts = prompts
self._loss_function = loss_function
self._main_device = main_device
self._compute_loss_in_worker = type(loss_function) is CrossEntropyLoss
self._use_prefix_cache = use_prefix_cache

def evaluate_candidates(
self,
Expand All @@ -86,11 +92,10 @@ def evaluate_candidates(

For each candidate group:
- Iterates through prompts sequentially.
- Dispatches ModelWorkerOperation.LOGITS with return_ids=True to all workers in parallel.
- Collects (logits, token_ids) from worker result queues.
- Computes loss using loss_function for target and control slices.
- Computes the built-in cross-entropy loss inside each model worker so full logits stay local.
- Preserves full-logits dispatch and parent-side computation for custom loss implementations.
- Accumulates losses on main_device.
- Releases intermediate logits and token_ids immediately to bound VRAM.
- Releases intermediate tensors after each prompt to bound VRAM.

Args:
control_candidates_by_group: List of candidate string lists per gradient shape group.
Expand Down Expand Up @@ -121,20 +126,36 @@ def evaluate_candidates(
prompt_indices = progress if progress is not None else range(num_prompts)

for i in prompt_indices:
for k, worker in enumerate(self._workers):
worker(self._prompts[k][i], ModelWorkerOperation.LOGITS, cand, return_ids=True)

logits, ids = zip(*[worker.results.get() for worker in self._workers], strict=True)
loss[j * batch_size : (j + 1) * batch_size] += sum(
self._loss_function.compute_loss(
logits=logit,
token_ids=token_ids,
target_slice=self._prompts[k][i]._target_slice,
control_slice=self._prompts[k][i]._control_slice,
).to(self._main_device)
for k, (logit, token_ids) in enumerate(zip(logits, ids, strict=True))
)
del logits, ids
if self._compute_loss_in_worker:
for k, worker in enumerate(self._workers):
worker(
self._prompts[k][i],
ModelWorkerOperation.LOSS,
cand,
self._loss_function,
use_prefix_cache=self._use_prefix_cache,
)

worker_losses = [worker.results.get() for worker in self._workers]
loss[j * batch_size : (j + 1) * batch_size] += sum(
worker_loss.to(self._main_device) for worker_loss in worker_losses
)
del worker_losses
else:
for k, worker in enumerate(self._workers):
worker(self._prompts[k][i], ModelWorkerOperation.LOGITS, cand, return_ids=True)

logits, ids = zip(*[worker.results.get() for worker in self._workers], strict=True)
loss[j * batch_size : (j + 1) * batch_size] += sum(
self._loss_function.compute_loss(
logits=logit,
token_ids=token_ids,
target_slice=self._prompts[k][i]._target_slice,
control_slice=self._prompts[k][i]._control_slice,
).to(self._main_device)
for k, (logit, token_ids) in enumerate(zip(logits, ids, strict=True))
)
del logits, ids

if progress is not None:
progress.set_description(
Expand Down
Loading
Loading