From db9454ad940b42acd5221649dd36c55b788bb4e8 Mon Sep 17 00:00:00 2001 From: Jayanth Sai Yarlagadda Date: Sat, 26 Sep 2026 16:16:03 -0700 Subject: [PATCH 1/3] PERF: Accelerate GCG candidate evaluation --- .../gcg/attack/base/attack_manager.py | 194 +++++++++++++++- .../gcg/attack/gcg/candidate_evaluator.py | 57 +++-- .../promptgen/gcg/attack/gcg/gcg_attack.py | 3 + pyrit/executor/promptgen/gcg/config.py | 4 + .../promptgen/gcg/default_implementations.py | 96 ++++++++ pyrit/executor/promptgen/gcg/generator.py | 1 + .../executor/promptgen/gcg/test_config.py | 9 +- .../gcg/test_default_implementations.py | 62 +++++ .../executor/promptgen/gcg/test_gcg_core.py | 55 ++++- .../promptgen/gcg/test_gcg_evaluation.py | 212 ++++++++++++++++++ .../executor/promptgen/gcg/test_generator.py | 2 + .../promptgen/gcg/trajectory_stubs.py | 8 + 12 files changed, 669 insertions(+), 34 deletions(-) diff --git a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py index b08b55c4cd..d2fa472a91 100644 --- a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py +++ b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py @@ -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 @@ -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] = [ @@ -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. @@ -610,15 +618,184 @@ 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. + """ + if hasattr(prefix_cache, "layers"): + expanded_cache = copy(prefix_cache) + expanded_layers = [] + for layer in prefix_cache.layers: + expanded_layer = copy(layer) + for attribute in ("keys", "values"): + value = getattr(layer, attribute, None) + if isinstance(value, torch.Tensor): + if value.shape[0] != 1: + raise ValueError("Prefix cache must be computed for exactly one sequence") + setattr(expanded_layer, attribute, value.expand(batch_size, *value.shape[1:])) + expanded_layers.append(expanded_layer) + expanded_cache.layers = expanded_layers + return expanded_cache + + if isinstance(prefix_cache, (tuple, list)): + return tuple(tuple(value.expand(batch_size, *value.shape[1:]) for value in layer) for layer in prefix_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: + """ + Score candidates after evaluating their invariant prefix once. + + Returns: + torch.Tensor: One scalar loss per candidate control. + """ + 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, + } + if attention_mask is not None: + prefix_kwargs["attention_mask"] = attention_mask[:1, :prefix_length] + prefix_output = model(**prefix_kwargs) + prefix_cache = self._expand_prefix_cache(prefix_output.past_key_values, token_ids.shape[0]) + + 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_output, 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: + return self._loss_with_prefix_cache( + model, + test_controls, + selective_loss, + logit_positions, + ) + 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. @@ -2051,6 +2228,7 @@ class ModelWorkerOperation(str, Enum): GRAD = "grad" LOGITS = "logits" + LOSS = "loss" CONTRAST_LOGITS = "contrast_logits" TEST = "test" TEST_LOSS = "test_loss" diff --git a/pyrit/executor/promptgen/gcg/attack/gcg/candidate_evaluator.py b/pyrit/executor/promptgen/gcg/attack/gcg/candidate_evaluator.py index 343947dd5c..20bf3ec571 100644 --- a/pyrit/executor/promptgen/gcg/attack/gcg/candidate_evaluator.py +++ b/pyrit/executor/promptgen/gcg/attack/gcg/candidate_evaluator.py @@ -13,6 +13,7 @@ ModelWorkerOperation, PromptManager, ) +from pyrit.executor.promptgen.gcg.default_implementations import CrossEntropyLoss from pyrit.executor.promptgen.gcg.extension_protocols import LossFunction @@ -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. @@ -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, @@ -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, @@ -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. @@ -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( diff --git a/pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py b/pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py index 46a4ff24b5..db4e341055 100644 --- a/pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py +++ b/pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py @@ -156,6 +156,7 @@ def __init__( sampling: SamplingStrategy | None = None, loss: LossFunction | None = None, candidate_filter: CandidateFilter | None = None, + use_prefix_cache: bool = False, ) -> None: """Initialize a GCG attack with optional algorithm extensions.""" super().__init__( @@ -173,6 +174,7 @@ def __init__( self._sampling = sampling self._loss = loss self._candidate_filter = candidate_filter + self._use_prefix_cache = use_prefix_cache def _resolve_sampling(self) -> SamplingStrategy: sampling: SamplingStrategy | None = getattr(self, "_sampling", None) @@ -332,6 +334,7 @@ def step( prompts=self.prompts, loss_function=loss_function, main_device=main_device, + use_prefix_cache=getattr(self, "_use_prefix_cache", False), ) eval_batch = evaluator.evaluate_candidates( control_candidates_by_group=candidate_batch.control_candidates_by_group, diff --git a/pyrit/executor/promptgen/gcg/config.py b/pyrit/executor/promptgen/gcg/config.py index d215ffd356..483f6b389f 100644 --- a/pyrit/executor/promptgen/gcg/config.py +++ b/pyrit/executor/promptgen/gcg/config.py @@ -185,6 +185,9 @@ class GCGAlgorithmConfig: Defaults to False. filter_cand (bool): Drop candidates whose token-length changes after re-tokenization. Defaults to True. + use_prefix_cache (bool): Cache the invariant prompt prefix while + scoring candidate suffixes. Defaults to False because reduced- + precision cache reuse can introduce small numerical differences. random_seed (int): Seed for ``torch``/``numpy``/``random``. Defaults to 42. control_init (str): Initial suffix string the optimization starts from. Defaults to twenty space-separated ``!`` tokens. @@ -212,6 +215,7 @@ class GCGAlgorithmConfig: learning_rate: float = 0.01 allow_non_ascii: bool = False filter_cand: bool = True + use_prefix_cache: bool = False random_seed: int = 42 control_init: str = _DEFAULT_CONTROL_INIT sampling: SamplingStrategy | None = None diff --git a/pyrit/executor/promptgen/gcg/default_implementations.py b/pyrit/executor/promptgen/gcg/default_implementations.py index 2686d296b4..1c3a857a5f 100644 --- a/pyrit/executor/promptgen/gcg/default_implementations.py +++ b/pyrit/executor/promptgen/gcg/default_implementations.py @@ -214,6 +214,102 @@ def compute_loss( result: torch.Tensor = total return result + def get_required_logit_positions( + self, + *, + target_slice: slice, + control_slice: slice, + device: torch.device | str | None = None, + ) -> torch.Tensor: + """ + Return sequence positions whose logits contribute to this loss. + + Causal language-model logits at position ``n - 1`` predict the token + at position ``n``. Disabled loss terms are omitted entirely so models + supporting ``logits_to_keep`` need not project unused hidden states + through the vocabulary-sized LM head. + + Args: + target_slice (slice): Target-token positions in the full sequence. + control_slice (slice): Control-token positions in the full sequence. + device (torch.device | str | None): Device on which to create the + returned index tensor. + + Returns: + torch.Tensor: Ordered sequence positions required by the enabled + target and control loss terms. + """ + positions: list[torch.Tensor] = [] + if self._target_weight > 0: + positions.append(torch.arange(target_slice.start - 1, target_slice.stop - 1, device=device)) + if self._control_weight > 0: + positions.append(torch.arange(control_slice.start - 1, control_slice.stop - 1, device=device)) + return torch.cat(positions) + + def compute_loss_from_selected_logits( + self, + *, + logits: torch.Tensor, + token_ids: torch.Tensor, + target_slice: slice, + control_slice: slice, + ) -> torch.Tensor: + """ + Compute loss from logits ordered by ``get_required_logit_positions``. + + Args: + logits (torch.Tensor): Selected logits with shape + ``(batch_size, required_positions, vocab_size)``. + token_ids (torch.Tensor): Full input token ids with shape + ``(batch_size, sequence_length)``. + target_slice (slice): Target-token positions in ``token_ids``. + control_slice (slice): Control-token positions in ``token_ids``. + + Returns: + torch.Tensor: Per-candidate scalar loss with shape ``(batch_size,)``. + + Raises: + ValueError: If the selected-logit count does not match the enabled + target and control terms. + RuntimeError: If both loss terms are unexpectedly disabled. + """ + target_length = target_slice.stop - target_slice.start if self._target_weight > 0 else 0 + control_length = control_slice.stop - control_slice.start if self._control_weight > 0 else 0 + expected_length = target_length + control_length + if logits.shape[1] != expected_length: + raise ValueError( + "Selected logits must contain one position per enabled loss token; " + f"expected {expected_length}, got {logits.shape[1]}" + ) + + criterion = nn.CrossEntropyLoss(reduction="none") + total: torch.Tensor | None = None + offset = 0 + + if self._target_weight > 0: + target_term = criterion( + logits[:, offset : offset + target_length, :].transpose(1, 2), + token_ids[:, target_slice], + ).mean(dim=-1) + total = self._target_weight * target_term + offset += target_length + + if self._control_weight > 0: + control_term = criterion( + logits[:, offset : offset + control_length, :].transpose(1, 2), + token_ids[:, control_slice], + ).mean(dim=-1) + weighted_control = self._control_weight * control_term + total = weighted_control if total is None else total + weighted_control + + if total is None: + raise RuntimeError( + "CrossEntropyLoss.compute_loss_from_selected_logits produced no terms; " + "this indicates a corrupted instance with both weights at 0." + ) + result: torch.Tensor = total + return result + class LengthPreservingFilter: """ diff --git a/pyrit/executor/promptgen/gcg/generator.py b/pyrit/executor/promptgen/gcg/generator.py index 119d59382e..2193928ca1 100644 --- a/pyrit/executor/promptgen/gcg/generator.py +++ b/pyrit/executor/promptgen/gcg/generator.py @@ -296,6 +296,7 @@ async def _perform_async(self, *, context: GCGContext) -> GCGResult: sampling=self._algorithm.sampling, loss=self._algorithm.loss, candidate_filter=self._algorithm.candidate_filter, + use_prefix_cache=self._algorithm.use_prefix_cache, ), } context.attack = self._create_attack( diff --git a/tests/unit/executor/promptgen/gcg/test_config.py b/tests/unit/executor/promptgen/gcg/test_config.py index 09665f9733..abae513cab 100644 --- a/tests/unit/executor/promptgen/gcg/test_config.py +++ b/tests/unit/executor/promptgen/gcg/test_config.py @@ -89,6 +89,7 @@ def test_minimal_config_constructs_with_defaults() -> None: assert config.algorithm.loss is None assert config.algorithm.candidate_filter is None assert config.algorithm.suffix_init is None + assert config.algorithm.use_prefix_cache is False assert config.strategy.transfer is False assert config.output.verbose is True assert config.hf_token is None @@ -211,7 +212,13 @@ def test_to_json_round_trip_preserves_all_fields() -> None: GCGModelConfig(name="mistralai/Mistral-7B-Instruct-v0.2"), ], test_models=[GCGModelConfig(name="lmsys/vicuna-7b-v1.5")], - algorithm=GCGAlgorithmConfig(n_steps=42, batch_size=64, target_weight=0.5, control_weight=0.5), + algorithm=GCGAlgorithmConfig( + n_steps=42, + batch_size=64, + target_weight=0.5, + control_weight=0.5, + use_prefix_cache=True, + ), strategy=GCGStrategyConfig(transfer=True, progressive_goals=True, anneal=True), output=GCGOutputConfig(result_prefix="results/run1", verbose=False), hf_token="hf_secrettoken", diff --git a/tests/unit/executor/promptgen/gcg/test_default_implementations.py b/tests/unit/executor/promptgen/gcg/test_default_implementations.py index 10f2474dc2..f9cc3218c2 100644 --- a/tests/unit/executor/promptgen/gcg/test_default_implementations.py +++ b/tests/unit/executor/promptgen/gcg/test_default_implementations.py @@ -306,6 +306,68 @@ def test_compute_loss_returns_batch_sized_tensor(self) -> None: assert out.shape == (batch_size,) + @pytest.mark.parametrize( + ("target_weight", "control_weight"), + [(1.0, 0.0), (0.0, 0.5), (0.7, 0.3)], + ) + def test_selected_logits_match_full_logits(self, target_weight: float, control_weight: float) -> None: + batch_size = 3 + target_slice = slice(5, 8) + control_slice = slice(1, 4) + torch.manual_seed(21) + logits = torch.randn(batch_size, 10, 25) + token_ids = torch.randint(0, 25, (batch_size, 10)) + loss_function = CrossEntropyLoss(target_weight=target_weight, control_weight=control_weight) + + positions = loss_function.get_required_logit_positions( + target_slice=target_slice, + control_slice=control_slice, + ) + selected = loss_function.compute_loss_from_selected_logits( + logits=logits[:, positions, :], + token_ids=token_ids, + target_slice=target_slice, + control_slice=control_slice, + ) + full = loss_function.compute_loss( + logits=logits, + token_ids=token_ids, + target_slice=target_slice, + control_slice=control_slice, + ) + + assert torch.equal(selected, full) + + def test_selected_logit_positions_follow_enabled_terms(self) -> None: + target_only = CrossEntropyLoss(target_weight=1.0, control_weight=0.0) + combined = CrossEntropyLoss(target_weight=1.0, control_weight=0.1) + + assert torch.equal( + target_only.get_required_logit_positions( + target_slice=slice(5, 8), + control_slice=slice(1, 4), + ), + torch.tensor([4, 5, 6]), + ) + assert torch.equal( + combined.get_required_logit_positions( + target_slice=slice(5, 8), + control_slice=slice(1, 4), + ), + torch.tensor([4, 5, 6, 0, 1, 2]), + ) + + def test_selected_logits_reject_wrong_sequence_length(self) -> None: + loss_function = CrossEntropyLoss(target_weight=1.0, control_weight=0.1) + + with pytest.raises(ValueError, match="expected 6, got 5"): + loss_function.compute_loss_from_selected_logits( + logits=torch.randn(2, 5, 20), + token_ids=torch.randint(0, 20, (2, 10)), + target_slice=slice(5, 8), + control_slice=slice(1, 4), + ) + def _make_filter_tokenizer() -> MagicMock: """Build a fresh, deterministic, stateless mock tokenizer for filter tests. diff --git a/tests/unit/executor/promptgen/gcg/test_gcg_core.py b/tests/unit/executor/promptgen/gcg/test_gcg_core.py index c9983d65de..6056b22bce 100644 --- a/tests/unit/executor/promptgen/gcg/test_gcg_core.py +++ b/tests/unit/executor/promptgen/gcg/test_gcg_core.py @@ -46,6 +46,7 @@ ) LengthPreservingFilter = default_implementations_mod.LengthPreservingFilter StandardGCGSampling = default_implementations_mod.StandardGCGSampling +CrossEntropyLoss = default_implementations_mod.CrossEntropyLoss import numpy as np # noqa: E402 @@ -1141,6 +1142,7 @@ def test_model_worker_task_payload_excludes_model() -> None: [ (ModelWorkerOperation.GRAD, "grad"), (ModelWorkerOperation.LOGITS, "logits"), + (ModelWorkerOperation.LOSS, "loss"), (ModelWorkerOperation.CONTRAST_LOGITS, "contrast_logits"), (ModelWorkerOperation.TEST, "test"), (ModelWorkerOperation.TEST_LOSS, "test_loss"), @@ -1191,6 +1193,9 @@ def __init__(self, items: list[Any]) -> None: def get(self) -> Any: return self._items.pop(0) + def put(self, item: Any) -> None: + self._items.append(item) + class _WorkerStub: def __init__( @@ -1204,11 +1209,29 @@ def __init__( self.model = MagicMock() self.model.device = "cpu" self.tokenizer = tokenizer - self.results = _Queue([gradient, (logits, token_ids)]) + self.results = _Queue([]) + self._gradient = gradient + self._logits = logits + self._token_ids = token_ids self.calls: list[tuple] = [] def __call__(self, *args: Any, **kwargs: Any) -> None: self.calls.append((args, kwargs)) + prompt, operation, *operation_args = args + if operation is ModelWorkerOperation.GRAD: + self.results.put(self._gradient) + elif operation is ModelWorkerOperation.LOGITS: + self.results.put((self._logits, self._token_ids)) + elif operation is ModelWorkerOperation.LOSS: + loss_function = operation_args[1] + self.results.put( + loss_function.compute_loss( + logits=self._logits, + token_ids=self._token_ids, + target_slice=prompt._target_slice, + control_slice=prompt._control_slice, + ) + ) class _PromptManagerStub: @@ -1441,12 +1464,13 @@ def test_step_default_path_matches_legacy_behavior(self) -> None: grad_args, grad_kwargs = worker.calls[0] assert grad_args == (prompt_manager, ModelWorkerOperation.GRAD) assert grad_kwargs == {} - logits_args, logits_kwargs = worker.calls[1] - assert logits_args[0] is prompt - assert logits_args[1] is ModelWorkerOperation.LOGITS - assert len(logits_args) == 3 - assert all(argument is not worker.model for argument in logits_args) - assert logits_kwargs == {"return_ids": True} + loss_args, loss_kwargs = worker.calls[1] + assert loss_args[0] is prompt + assert loss_args[1] is ModelWorkerOperation.LOSS + assert loss_args[2] == legacy_controls + assert isinstance(loss_args[3], CrossEntropyLoss) + assert all(argument is not worker.model for argument in loss_args) + assert loss_kwargs == {"use_prefix_cache": False} def test_step_uses_custom_protocol_implementations_when_supplied(self) -> None: gradient = torch.randn(3, 6) @@ -1684,6 +1708,23 @@ def test_attack_prompt_logits_builds_attention_mask() -> None: assert torch.equal(model.call_args.kwargs["attention_mask"], torch.ones(1, 4, dtype=torch.long)) +def test_attack_prompt_logits_forwards_selected_positions() -> None: + prompt = object.__new__(AttackPrompt) + prompt._control_slice = slice(1, 3) + prompt.input_ids = torch.tensor([0, 1, 2, 3]) + prompt.tokenizer = MagicMock() + prompt.tokenizer.return_value.input_ids = [5, 6] + model = MagicMock() + model.device = torch.device("cpu") + model.return_value.logits = torch.randn(1, 2, 8) + positions = torch.tensor([0, 2]) + + logits = prompt.logits(model, test_controls=["candidate"], logits_to_keep=positions) + + assert logits.shape == (1, 2, 8) + assert torch.equal(model.call_args.kwargs["logits_to_keep"], positions) + + def test_prompt_manager_grad_streams_and_sums_prompt_gradients() -> None: prompt_manager = object.__new__(PromptManager) first_prompt = MagicMock() diff --git a/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py b/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py index 909d3cf757..2cca8ec7d7 100644 --- a/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py +++ b/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock import pytest +from transformers import Qwen2Config, Qwen2ForCausalLM pytest.importorskip( "pyrit.executor.promptgen.gcg.attack.base.attack_manager", @@ -23,6 +24,7 @@ CandidateEvaluationBatch, GCGCandidateEvaluator, ) +from pyrit.executor.promptgen.gcg.default_implementations import CrossEntropyLoss from pyrit.executor.promptgen.gcg.extension_protocols import LossFunction @@ -191,6 +193,216 @@ def test_group_candidate_count_mismatch_with_batch_size_raises(self) -> None: class TestGCGCandidateEvaluatorExecution: + def test_selective_logits_match_full_logits_on_transformers_model(self) -> None: + torch.manual_seed(123) + model = Qwen2ForCausalLM( + Qwen2Config( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=2, + max_position_embeddings=32, + ) + ).eval() + prompt = object.__new__(AttackPrompt) + prompt._control_slice = slice(1, 3) + prompt._target_slice = slice(4, 7) + prompt.input_ids = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8]) + prompt.tokenizer = MagicMock() + candidates = torch.tensor([[9, 10], [11, 12]]) + loss_fn = CrossEntropyLoss(target_weight=0.7, control_weight=0.3) + + full_logits, token_ids = prompt.logits(model, candidates, return_ids=True) + expected = loss_fn.compute_loss( + logits=full_logits, + token_ids=token_ids, + target_slice=prompt._target_slice, + control_slice=prompt._control_slice, + ) + actual = prompt.loss(model, candidates, loss_fn) + + assert torch.equal(actual, expected) + + def test_prefix_cached_loss_matches_full_forward_on_transformers_model(self) -> None: + torch.manual_seed(123) + model = Qwen2ForCausalLM( + Qwen2Config( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + max_position_embeddings=32, + ) + ).eval() + prompt = object.__new__(AttackPrompt) + prompt._control_slice = slice(3, 5) + prompt._target_slice = slice(6, 9) + prompt.input_ids = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) + prompt.tokenizer = MagicMock() + candidates = torch.tensor([[11, 12], [13, 14], [15, 16]]) + loss_fn = CrossEntropyLoss(target_weight=0.7, control_weight=0.3) + + full_logits, token_ids = prompt.logits(model, candidates, return_ids=True) + expected = loss_fn.compute_loss( + logits=full_logits, + token_ids=token_ids, + target_slice=prompt._target_slice, + control_slice=prompt._control_slice, + ) + actual = prompt.loss(model, candidates, loss_fn, use_prefix_cache=True) + + assert torch.allclose(actual, expected, rtol=1e-5, atol=1e-6) + assert actual.argmin() == expected.argmin() + + def test_attack_prompt_worker_loss_matches_direct_computation(self) -> None: + logits = torch.randn(2, 6, 10) + token_ids = torch.randint(0, 10, (2, 6)) + prompt = MagicMock() + prompt.logits.return_value = (logits, token_ids) + prompt._target_slice = slice(3, 5) + prompt._control_slice = slice(1, 3) + loss_fn = CrossEntropyLoss(target_weight=0.7, control_weight=0.3) + model = MagicMock() + candidates = ["cand-1", "cand-2"] + + actual = AttackPrompt.loss(prompt, model, candidates, loss_fn) + expected = loss_fn.compute_loss( + logits=logits, + token_ids=token_ids, + target_slice=prompt._target_slice, + control_slice=prompt._control_slice, + ) + + assert torch.equal(actual, expected) + prompt.logits.assert_called_once_with(model, candidates, return_ids=True) + + def test_attack_prompt_worker_loss_selects_only_required_logits_when_supported(self) -> None: + class SelectiveModel: + device = torch.device("cpu") + + def forward(self, *, logits_to_keep: int | torch.Tensor = 0) -> None: + del logits_to_keep + + full_logits = torch.randn(2, 6, 10) + token_ids = torch.randint(0, 10, (2, 6)) + prompt = MagicMock() + prompt._target_slice = slice(3, 5) + prompt._control_slice = slice(1, 3) + loss_fn = CrossEntropyLoss(target_weight=0.7, control_weight=0.3) + expected_positions = torch.tensor([2, 3, 0, 1]) + prompt.logits.return_value = (full_logits[:, expected_positions, :], token_ids) + model = SelectiveModel() + + actual = AttackPrompt.loss(prompt, model, ["cand-1", "cand-2"], loss_fn) + expected = loss_fn.compute_loss( + logits=full_logits, + token_ids=token_ids, + target_slice=prompt._target_slice, + control_slice=prompt._control_slice, + ) + + assert torch.equal(actual, expected) + prompt.logits.assert_called_once() + args, kwargs = prompt.logits.call_args + assert args == (model, ["cand-1", "cand-2"]) + assert kwargs["return_ids"] is True + assert torch.equal(kwargs["logits_to_keep"], expected_positions) + + def test_builtin_loss_is_computed_in_worker(self) -> None: + prompt = _MockAttackPrompt(target_slice=slice(3, 5), control_slice=slice(1, 3)) + pm = _MockPromptManager([prompt]) # type: ignore[arg-type] + computed_losses = torch.tensor([0.42, 0.99]) + worker = _MockWorker([computed_losses]) + loss_fn = CrossEntropyLoss(target_weight=1.0, control_weight=0.1) + evaluator = GCGCandidateEvaluator( + workers=[worker], # type: ignore[list-item] + prompts=[pm], # type: ignore[list-item] + loss_function=loss_fn, + main_device=torch.device("cpu"), + ) + + candidates = [["cand-1", "cand-2"]] + result = evaluator.evaluate_candidates(control_candidates_by_group=candidates, batch_size=2) + + assert torch.equal(result.losses, computed_losses) + assert len(worker.calls) == 1 + args, kwargs = worker.calls[0] + assert args == (prompt, ModelWorkerOperation.LOSS, candidates[0], loss_fn) + assert kwargs == {"use_prefix_cache": False} + + def test_builtin_loss_accumulates_across_workers(self) -> None: + prompts = [ + _MockPromptManager([_MockAttackPrompt(slice(3, 5), slice(1, 3))]), + _MockPromptManager([_MockAttackPrompt(slice(3, 5), slice(1, 3))]), + ] + workers = [ + _MockWorker([torch.tensor([0.2, 0.3])]), + _MockWorker([torch.tensor([0.5, 0.7])]), + ] + loss_fn = CrossEntropyLoss() + evaluator = GCGCandidateEvaluator( + workers=workers, # type: ignore[arg-type] + prompts=prompts, # type: ignore[arg-type] + loss_function=loss_fn, + main_device=torch.device("cpu"), + ) + + result = evaluator.evaluate_candidates( + control_candidates_by_group=[["cand-1", "cand-2"]], + batch_size=2, + ) + + assert torch.allclose(result.losses, torch.tensor([0.7, 1.0])) + assert all(worker.calls[0][0][1] is ModelWorkerOperation.LOSS for worker in workers) + + def test_prefix_cache_opt_in_is_forwarded_to_worker(self) -> None: + prompt = _MockAttackPrompt(target_slice=slice(3, 5), control_slice=slice(1, 3)) + pm = _MockPromptManager([prompt]) # type: ignore[arg-type] + computed_losses = torch.tensor([0.42, 0.99]) + worker = _MockWorker([computed_losses]) + loss_fn = CrossEntropyLoss() + evaluator = GCGCandidateEvaluator( + workers=[worker], # type: ignore[list-item] + prompts=[pm], # type: ignore[list-item] + loss_function=loss_fn, + main_device=torch.device("cpu"), + use_prefix_cache=True, + ) + + evaluator.evaluate_candidates(control_candidates_by_group=[["cand-1", "cand-2"]], batch_size=2) + + assert worker.calls == [ + ( + (prompt, ModelWorkerOperation.LOSS, ["cand-1", "cand-2"], loss_fn), + {"use_prefix_cache": True}, + ) + ] + + def test_builtin_loss_subclass_uses_custom_loss_path(self) -> None: + class CustomCrossEntropyLoss(CrossEntropyLoss): + pass + + prompt = _MockAttackPrompt(target_slice=slice(3, 5), control_slice=slice(1, 3)) + pm = _MockPromptManager([prompt]) # type: ignore[arg-type] + logits = torch.randn(2, 6, 10) + token_ids = torch.randint(0, 10, (2, 6)) + worker = _MockWorker([(logits, token_ids)]) + loss_fn = CustomCrossEntropyLoss() + evaluator = GCGCandidateEvaluator( + workers=[worker], # type: ignore[list-item] + prompts=[pm], # type: ignore[list-item] + loss_function=loss_fn, + main_device=torch.device("cpu"), + ) + + evaluator.evaluate_candidates(control_candidates_by_group=[["cand-1", "cand-2"]], batch_size=2) + + assert worker.calls[0][0][1] is ModelWorkerOperation.LOGITS + def test_single_worker_single_prompt_evaluation(self) -> None: target_slice = slice(3, 5) control_slice = slice(1, 3) diff --git a/tests/unit/executor/promptgen/gcg/test_generator.py b/tests/unit/executor/promptgen/gcg/test_generator.py index 72923923d3..a9ef9500d7 100644 --- a/tests/unit/executor/promptgen/gcg/test_generator.py +++ b/tests/unit/executor/promptgen/gcg/test_generator.py @@ -369,6 +369,7 @@ def filter_candidates( sampling=sampling, loss=loss, candidate_filter=candidate_filter, + use_prefix_cache=True, ), output=GCGOutputConfig(result_prefix=str(tmp_path / "gcg")), ) @@ -395,6 +396,7 @@ def filter_candidates( assert mpa_factory.keywords["sampling"] is sampling assert mpa_factory.keywords["loss"] is loss assert mpa_factory.keywords["candidate_filter"] is candidate_filter + assert mpa_factory.keywords["use_prefix_cache"] is True class TestReadResult: diff --git a/tests/unit/executor/promptgen/gcg/trajectory_stubs.py b/tests/unit/executor/promptgen/gcg/trajectory_stubs.py index 0fcc1b052c..c7cafad97a 100644 --- a/tests/unit/executor/promptgen/gcg/trajectory_stubs.py +++ b/tests/unit/executor/promptgen/gcg/trajectory_stubs.py @@ -147,6 +147,14 @@ def _execute(self, ob: Any, operation: ModelWorkerOperation, *args: Any) -> Any: return self._grad(ob.control_toks) if operation is ModelWorkerOperation.LOGITS: return self._logits(ob, args[0]) + if operation is ModelWorkerOperation.LOSS: + logits, token_ids = self._logits(ob, args[0]) + return args[1].compute_loss( + logits=logits, + token_ids=token_ids, + target_slice=ob._target_slice, + control_slice=ob._control_slice, + ) if operation is ModelWorkerOperation.TEST: return [(ob.control_str != ob.control_init, 0) for _ in ob] if operation is ModelWorkerOperation.TEST_LOSS: From 01b092f30c5d18873c5a47506b0f8d7801311f5c Mon Sep 17 00:00:00 2001 From: Jayanth Sai Yarlagadda Date: Sun, 27 Sep 2026 21:39:55 -0700 Subject: [PATCH 2/3] FIX: Fall back for unsupported GCG prefix caches --- .../gcg/attack/base/attack_manager.py | 61 +++++++++++--- .../promptgen/gcg/test_gcg_evaluation.py | 80 ++++++++++++++++++- 2 files changed, 127 insertions(+), 14 deletions(-) diff --git a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py index d2fa472a91..1b46944e7f 100644 --- a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py +++ b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py @@ -661,23 +661,49 @@ def _expand_prefix_cache(prefix_cache: Any, batch_size: int) -> Any: 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) - for attribute in ("keys", "values"): - value = getattr(layer, attribute, None) - if isinstance(value, torch.Tensor): - if value.shape[0] != 1: - raise ValueError("Prefix cache must be computed for exactly one sequence") - setattr(expanded_layer, attribute, value.expand(batch_size, *value.shape[1:])) + 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)): - return tuple(tuple(value.expand(batch_size, *value.shape[1:]) for value in layer) for layer in prefix_cache) + 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}") @@ -687,12 +713,13 @@ def _loss_with_prefix_cache( test_controls: Any, loss_function: Any, logit_positions: torch.Tensor, - ) -> torch.Tensor: + ) -> torch.Tensor | None: """ Score candidates after evaluating their invariant prefix once. Returns: - torch.Tensor: One scalar loss per candidate control. + 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 @@ -700,11 +727,19 @@ def _loss_with_prefix_cache( "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) - prefix_cache = self._expand_prefix_cache(prefix_output.past_key_values, token_ids.shape[0]) + 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:], @@ -716,7 +751,7 @@ def _loss_with_prefix_cache( if attention_mask is not None: suffix_kwargs["attention_mask"] = attention_mask logits = model(**suffix_kwargs).logits - del prefix_output, prefix_cache + del prefix_cache try: result: torch.Tensor = loss_function.compute_loss_from_selected_logits( @@ -762,12 +797,14 @@ def loss( device=model.device, ) if use_prefix_cache and supports_prefix_cache and self._control_slice.start > 1: - return self._loss_with_prefix_cache( + 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, diff --git a/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py b/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py index 2cca8ec7d7..566c28a264 100644 --- a/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py +++ b/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py @@ -7,7 +7,7 @@ from unittest.mock import MagicMock import pytest -from transformers import Qwen2Config, Qwen2ForCausalLM +from transformers import Qwen2Config, Qwen2ForCausalLM, Qwen3NextConfig, Qwen3NextForCausalLM # type: ignore[ty:possibly-missing-import] pytest.importorskip( "pyrit.executor.promptgen.gcg.attack.base.attack_manager", @@ -193,6 +193,19 @@ def test_group_candidate_count_mismatch_with_batch_size_raises(self) -> None: class TestGCGCandidateEvaluatorExecution: + def test_prefix_cache_rejects_additional_batched_layer_state(self) -> None: + class CacheLayer: + def __init__(self) -> None: + self.keys = torch.zeros(1, 2, 3, 4) + self.values = torch.zeros(1, 2, 3, 4) + self.conv_states = {0: torch.zeros(1, 2, 3)} + + class Cache: + layers = [CacheLayer()] + + with pytest.raises(TypeError, match="Unsupported state"): + AttackPrompt._expand_prefix_cache(Cache(), batch_size=2) + def test_selective_logits_match_full_logits_on_transformers_model(self) -> None: torch.manual_seed(123) model = Qwen2ForCausalLM( @@ -253,10 +266,73 @@ def test_prefix_cached_loss_matches_full_forward_on_transformers_model(self) -> target_slice=prompt._target_slice, control_slice=prompt._control_slice, ) - actual = prompt.loss(model, candidates, loss_fn, use_prefix_cache=True) + forward_calls = [] + + def capture_forward(_module, _args, kwargs, output): + forward_calls.append((kwargs, output.logits.shape)) + + hook = model.register_forward_hook(capture_forward, with_kwargs=True) + try: + actual = prompt.loss(model, candidates, loss_fn, use_prefix_cache=True) + finally: + hook.remove() assert torch.allclose(actual, expected, rtol=1e-5, atol=1e-6) assert actual.argmin() == expected.argmin() + prefix_kwargs, prefix_logits_shape = forward_calls[0] + assert prefix_kwargs["logits_to_keep"] == 1 + assert prefix_logits_shape[1] == 1 + + def test_hybrid_cache_falls_back_to_uncached_selective_logits(self) -> None: + torch.manual_seed(123) + model = Qwen3NextForCausalLM( + Qwen3NextConfig( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=1, + head_dim=8, + max_position_embeddings=32, + linear_conv_kernel_dim=4, + linear_key_head_dim=8, + linear_value_head_dim=8, + linear_num_key_heads=2, + linear_num_value_heads=2, + moe_intermediate_size=16, + shared_expert_intermediate_size=16, + num_experts_per_tok=1, + num_experts=2, + layer_types=["linear_attention", "full_attention"], + ) + ).eval() + prompt = object.__new__(AttackPrompt) + prompt._control_slice = slice(3, 5) + prompt._target_slice = slice(6, 9) + prompt.input_ids = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) + prompt.tokenizer = MagicMock() + candidates = torch.tensor([[11, 12], [13, 14], [15, 16]]) + loss_fn = CrossEntropyLoss(target_weight=0.7, control_weight=0.3) + + expected = prompt.loss(model, candidates, loss_fn) + forward_calls = [] + + def capture_forward(_module, _args, kwargs, _output): + forward_calls.append(kwargs) + + hook = model.register_forward_hook(capture_forward, with_kwargs=True) + try: + actual = prompt.loss(model, candidates, loss_fn, use_prefix_cache=True) + finally: + hook.remove() + + assert torch.equal(actual, expected) + assert len(forward_calls) == 2 + assert forward_calls[0]["input_ids"].shape[0] == 1 + assert forward_calls[0]["logits_to_keep"] == 1 + assert forward_calls[1]["input_ids"].shape[0] == len(candidates) + assert "past_key_values" not in forward_calls[1] def test_attack_prompt_worker_loss_matches_direct_computation(self) -> None: logits = torch.randn(2, 6, 10) From 3a4e7ca61605bd51ab0e2b1070a209fedadf9471 Mon Sep 17 00:00:00 2001 From: Jayanth Sai Yarlagadda Date: Mon, 28 Sep 2026 09:55:33 -0700 Subject: [PATCH 3/3] TEST: Cover GCG prefix-cache fallbacks --- .../gcg/test_default_implementations.py | 13 +++ .../promptgen/gcg/test_gcg_evaluation.py | 83 ++++++++++++++++++- 2 files changed, 95 insertions(+), 1 deletion(-) diff --git a/tests/unit/executor/promptgen/gcg/test_default_implementations.py b/tests/unit/executor/promptgen/gcg/test_default_implementations.py index f9cc3218c2..f3cc41cd13 100644 --- a/tests/unit/executor/promptgen/gcg/test_default_implementations.py +++ b/tests/unit/executor/promptgen/gcg/test_default_implementations.py @@ -368,6 +368,19 @@ def test_selected_logits_reject_wrong_sequence_length(self) -> None: control_slice=slice(1, 4), ) + def test_selected_logits_reject_corrupted_zero_weight_state(self) -> None: + loss_function = CrossEntropyLoss() + loss_function._target_weight = 0.0 + loss_function._control_weight = 0.0 + + with pytest.raises(RuntimeError, match="produced no terms"): + loss_function.compute_loss_from_selected_logits( + logits=torch.empty(2, 0, 20), + token_ids=torch.randint(0, 20, (2, 10)), + target_slice=slice(5, 8), + control_slice=slice(1, 4), + ) + def _make_filter_tokenizer() -> MagicMock: """Build a fresh, deterministic, stateless mock tokenizer for filter tests. diff --git a/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py b/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py index 566c28a264..992bd96cb9 100644 --- a/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py +++ b/tests/unit/executor/promptgen/gcg/test_gcg_evaluation.py @@ -198,7 +198,7 @@ class CacheLayer: def __init__(self) -> None: self.keys = torch.zeros(1, 2, 3, 4) self.values = torch.zeros(1, 2, 3, 4) - self.conv_states = {0: torch.zeros(1, 2, 3)} + self.conv_states = {0: [torch.zeros(1, 2, 3)]} class Cache: layers = [CacheLayer()] @@ -206,6 +206,58 @@ class Cache: with pytest.raises(TypeError, match="Unsupported state"): AttackPrompt._expand_prefix_cache(Cache(), batch_size=2) + def test_prefix_cache_rejects_layer_without_instance_state(self) -> None: + class CacheLayer: + __slots__ = ("keys", "values") + + def __init__(self) -> None: + self.keys = torch.zeros(1, 2, 3, 4) + self.values = torch.zeros(1, 2, 3, 4) + + cache = MagicMock() + cache.layers = [CacheLayer()] + + with pytest.raises(TypeError, match="Unsupported prefix-cache layer"): + AttackPrompt._expand_prefix_cache(cache, batch_size=2) + + def test_prefix_cache_rejects_non_singleton_batch(self) -> None: + layer = MagicMock() + layer.keys = torch.zeros(2, 2, 3, 4) + layer.values = torch.zeros(2, 2, 3, 4) + cache = MagicMock() + cache.layers = [layer] + + with pytest.raises(ValueError, match="exactly one sequence"): + AttackPrompt._expand_prefix_cache(cache, batch_size=2) + + def test_prefix_cache_expands_legacy_cache(self) -> None: + keys = torch.randn(1, 2, 3, 4) + values = torch.randn(1, 2, 3, 4) + + expanded = AttackPrompt._expand_prefix_cache(((keys, values),), batch_size=3) + + assert expanded[0][0].shape == (3, 2, 3, 4) + assert expanded[0][1].shape == (3, 2, 3, 4) + assert torch.equal(expanded[0][0][0], keys[0]) + assert torch.equal(expanded[0][1][0], values[0]) + + @pytest.mark.parametrize( + ("cache", "expected_error", "message"), + [ + (((torch.zeros(1, 2), "not-a-tensor"),), TypeError, "Unsupported prefix-cache layer"), + (((torch.zeros(2, 2), torch.zeros(2, 2)),), ValueError, "exactly one sequence"), + (object(), TypeError, "Unsupported prefix-cache type"), + ], + ) + def test_prefix_cache_rejects_unsupported_legacy_cache( + self, + cache: object, + expected_error: type[Exception], + message: str, + ) -> None: + with pytest.raises(expected_error, match=message): + AttackPrompt._expand_prefix_cache(cache, batch_size=2) + def test_selective_logits_match_full_logits_on_transformers_model(self) -> None: torch.manual_seed(123) model = Qwen2ForCausalLM( @@ -266,6 +318,8 @@ def test_prefix_cached_loss_matches_full_forward_on_transformers_model(self) -> target_slice=prompt._target_slice, control_slice=prompt._control_slice, ) + attention_mask = torch.ones_like(token_ids) + prompt._build_candidate_batch = MagicMock(return_value=(token_ids, attention_mask)) forward_calls = [] def capture_forward(_module, _args, kwargs, output): @@ -282,6 +336,11 @@ def capture_forward(_module, _args, kwargs, output): prefix_kwargs, prefix_logits_shape = forward_calls[0] assert prefix_kwargs["logits_to_keep"] == 1 assert prefix_logits_shape[1] == 1 + assert torch.equal( + prefix_kwargs["attention_mask"], + attention_mask[:1, : prompt._control_slice.start - 1], + ) + assert torch.equal(forward_calls[1][0]["attention_mask"], attention_mask) def test_hybrid_cache_falls_back_to_uncached_selective_logits(self) -> None: torch.manual_seed(123) @@ -356,6 +415,28 @@ def test_attack_prompt_worker_loss_matches_direct_computation(self) -> None: assert torch.equal(actual, expected) prompt.logits.assert_called_once_with(model, candidates, return_ids=True) + def test_attack_prompt_loss_handles_uninspectable_model_forward(self) -> None: + logits = torch.randn(2, 6, 10) + token_ids = torch.randint(0, 10, (2, 6)) + prompt = MagicMock() + prompt.logits.return_value = (logits, token_ids) + prompt._target_slice = slice(3, 5) + prompt._control_slice = slice(1, 3) + loss_fn = CrossEntropyLoss(target_weight=0.7, control_weight=0.3) + model = MagicMock() + model.forward = object() + + actual = AttackPrompt.loss(prompt, model, ["cand-1", "cand-2"], loss_fn) + + expected = loss_fn.compute_loss( + logits=logits, + token_ids=token_ids, + target_slice=prompt._target_slice, + control_slice=prompt._control_slice, + ) + assert torch.equal(actual, expected) + prompt.logits.assert_called_once_with(model, ["cand-1", "cand-2"], return_ids=True) + def test_attack_prompt_worker_loss_selects_only_required_logits_when_supported(self) -> None: class SelectiveModel: device = torch.device("cpu")