Skip to content

Commit d620b4d

Browse files
Kaap10romanlutzCopilot
authored
MAINT GCG: extract candidate proposal phase (#2665) (#2671)
Co-authored-by: Roman Lutz <romanlutz13@gmail.com> Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 2429881 commit d620b4d

3 files changed

Lines changed: 572 additions & 53 deletions

File tree

Lines changed: 261 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,261 @@
1+
# Copyright (c) Microsoft Corporation.
2+
# Licensed under the MIT license.
3+
4+
"""Candidate proposal logic for Greedy Coordinate Gradient (GCG) attacks."""
5+
6+
from collections.abc import Callable
7+
from dataclasses import dataclass
8+
9+
import torch
10+
11+
from pyrit.executor.promptgen.gcg.attack.base.attack_manager import (
12+
ModelWorker,
13+
ModelWorkerOperation,
14+
PromptManager,
15+
)
16+
from pyrit.executor.promptgen.gcg.extension_protocols import CandidateFilter, SamplingStrategy
17+
18+
19+
@dataclass(frozen=True, slots=True)
20+
class CandidateProposalBatch:
21+
"""
22+
Grouped candidate control tokens and filtered text strings across compatible workers.
23+
24+
Attributes:
25+
control_candidates_by_group: List of filtered candidate string lists, one per
26+
compatible gradient shape group.
27+
group_worker_indices: Index of the last worker in each contiguous compatible
28+
gradient shape group, used for sampling and filtering that group.
29+
"""
30+
31+
control_candidates_by_group: list[list[str]]
32+
group_worker_indices: list[int]
33+
34+
@property
35+
def num_groups(self) -> int:
36+
"""The number of candidate groups."""
37+
return len(self.control_candidates_by_group)
38+
39+
40+
class GCGCandidateProposer:
41+
"""Encapsulates gradient aggregation, token candidate sampling, and candidate filtering for GCG."""
42+
43+
def __init__(
44+
self,
45+
*,
46+
workers: list[ModelWorker],
47+
prompts: list[PromptManager],
48+
sampling: SamplingStrategy | None = None,
49+
candidate_filter: CandidateFilter | None = None,
50+
sample_fn: Callable[..., torch.Tensor] | None = None,
51+
filter_fn: Callable[..., list[str]] | None = None,
52+
main_device: torch.device,
53+
) -> None:
54+
"""
55+
Initialize the candidate proposer with workers, prompt managers, and extension protocols.
56+
57+
Args:
58+
workers: List of model workers participating in the attack.
59+
prompts: List of prompt managers associated with each worker.
60+
sampling: Sampling strategy protocol used to sample candidate token indices from gradients.
61+
candidate_filter: Candidate filter protocol used to filter and decode candidate tokens.
62+
sample_fn: Optional callable to sample candidates. If omitted, uses sampling protocol.
63+
filter_fn: Optional callable to filter candidates. If omitted, uses candidate_filter protocol.
64+
main_device: Target PyTorch device on which gradients are aggregated.
65+
66+
Raises:
67+
ValueError: If workers list is empty or if worker and prompt manager counts mismatch.
68+
"""
69+
if not workers:
70+
raise ValueError("GCG optimization requires at least one worker")
71+
if len(workers) != len(prompts):
72+
raise ValueError("Worker and PromptManager count mismatch")
73+
74+
self._workers = workers
75+
self._prompts = prompts
76+
self._sampling = sampling
77+
self._candidate_filter = candidate_filter
78+
self._sample_fn = sample_fn
79+
self._filter_fn = filter_fn
80+
self._main_device = main_device
81+
82+
def propose_candidates(
83+
self,
84+
*,
85+
batch_size: int = 1024,
86+
topk: int = 256,
87+
temp: float = 1.0,
88+
allow_non_ascii: bool = True,
89+
filter_cand: bool = True,
90+
current_control_str: str,
91+
) -> CandidateProposalBatch:
92+
"""
93+
Dispatch gradient operations, aggregate compatible shapes, and sample/filter candidate controls.
94+
95+
Args:
96+
batch_size: Number of candidate controls per batch. Defaults to 1024.
97+
topk: Number of top gradient positions to sample from. Defaults to 256.
98+
temp: Temperature for sampling. Kept for protocol compatibility. Defaults to 1.0.
99+
allow_non_ascii: Whether to allow non-ASCII tokens. Defaults to True.
100+
filter_cand: Whether to filter invalid candidates. Defaults to True.
101+
current_control_str: The current decoded control string used as a fallback by length filters.
102+
103+
Returns:
104+
CandidateProposalBatch containing filtered candidate strings per gradient shape group
105+
and corresponding group worker indices.
106+
107+
Raises:
108+
RuntimeError: If workers do not produce an aggregate gradient.
109+
"""
110+
# Dispatch gradient calculation to all workers
111+
for j, worker in enumerate(self._workers):
112+
worker(self._prompts[j], ModelWorkerOperation.GRAD)
113+
114+
control_cands: list[list[str]] = []
115+
group_worker_indices: list[int] = []
116+
grad: torch.Tensor | None = None
117+
118+
# Collect and aggregate gradients across workers
119+
for j, worker in enumerate(self._workers):
120+
new_grad: torch.Tensor = worker.results.get().to(self._main_device)
121+
new_grad = new_grad / new_grad.norm(dim=-1, keepdim=True)
122+
123+
if grad is None:
124+
grad = torch.zeros_like(new_grad)
125+
126+
if grad.shape != new_grad.shape:
127+
# Shape mismatch: finalize the preceding group
128+
with torch.no_grad():
129+
sampled = self._sample_group(
130+
worker_idx=j - 1,
131+
gradient=grad,
132+
batch_size=batch_size,
133+
topk=topk,
134+
temp=temp,
135+
allow_non_ascii=allow_non_ascii,
136+
)
137+
filtered = self._filter_group(
138+
worker_idx=j - 1,
139+
control_cand=sampled,
140+
filter_cand=filter_cand,
141+
current_control_str=current_control_str,
142+
)
143+
control_cands.append(filtered)
144+
group_worker_indices.append(j - 1)
145+
grad = new_grad
146+
else:
147+
grad += new_grad
148+
149+
if grad is None:
150+
raise RuntimeError("GCG workers did not produce an aggregate gradient")
151+
152+
# Finalize the last group
153+
last_worker_idx = len(self._workers) - 1
154+
with torch.no_grad():
155+
sampled = self._sample_group(
156+
worker_idx=last_worker_idx,
157+
gradient=grad,
158+
batch_size=batch_size,
159+
topk=topk,
160+
temp=temp,
161+
allow_non_ascii=allow_non_ascii,
162+
)
163+
filtered = self._filter_group(
164+
worker_idx=last_worker_idx,
165+
control_cand=sampled,
166+
filter_cand=filter_cand,
167+
current_control_str=current_control_str,
168+
)
169+
control_cands.append(filtered)
170+
group_worker_indices.append(last_worker_idx)
171+
172+
return CandidateProposalBatch(
173+
control_candidates_by_group=control_cands,
174+
group_worker_indices=group_worker_indices,
175+
)
176+
177+
def _sample_group(
178+
self,
179+
*,
180+
worker_idx: int,
181+
gradient: torch.Tensor,
182+
batch_size: int,
183+
topk: int,
184+
temp: float,
185+
allow_non_ascii: bool,
186+
) -> torch.Tensor:
187+
"""
188+
Sample candidate token indices for a specific worker's control slice.
189+
190+
Args:
191+
worker_idx: Index of the representative worker.
192+
gradient: Aggregated gradient tensor for this shape group.
193+
batch_size: Number of candidates to sample.
194+
topk: Top gradient coordinates to sample from.
195+
temp: Sampling temperature.
196+
allow_non_ascii: Whether non-ASCII tokens are permitted.
197+
198+
Returns:
199+
Tensor of sampled candidate token IDs.
200+
201+
Raises:
202+
ValueError: If neither sample_fn nor sampling strategy was provided.
203+
"""
204+
if self._sample_fn is not None:
205+
return self._sample_fn(
206+
worker_index=worker_idx,
207+
gradient=gradient,
208+
batch_size=batch_size,
209+
topk=topk,
210+
temp=temp,
211+
allow_non_ascii=allow_non_ascii,
212+
)
213+
if self._sampling is None:
214+
raise ValueError("SamplingStrategy or sample_fn must be provided")
215+
prompt_manager = self._prompts[worker_idx]
216+
return self._sampling.sample_candidates(
217+
gradient=gradient,
218+
control_tokens=prompt_manager.control_toks,
219+
batch_size=batch_size,
220+
top_k=topk,
221+
temperature=temp,
222+
allow_non_ascii=allow_non_ascii,
223+
non_ascii_tokens=prompt_manager.disallowed_toks,
224+
)
225+
226+
def _filter_group(
227+
self,
228+
*,
229+
worker_idx: int,
230+
control_cand: torch.Tensor,
231+
filter_cand: bool,
232+
current_control_str: str,
233+
) -> list[str]:
234+
"""
235+
Filter and decode candidate token tensors into valid string controls.
236+
237+
Args:
238+
worker_idx: Index of the representative worker.
239+
control_cand: Sampled candidate token IDs tensor.
240+
filter_cand: Whether candidate filtering is enabled.
241+
current_control_str: Current decoded control string fallback.
242+
243+
Returns:
244+
List of decoded and filtered candidate control strings.
245+
246+
Raises:
247+
ValueError: If neither filter_fn nor candidate_filter was provided.
248+
"""
249+
if self._filter_fn is not None:
250+
return self._filter_fn(
251+
worker_index=worker_idx,
252+
control_cand=control_cand,
253+
filter_cand=filter_cand,
254+
)
255+
if self._candidate_filter is None:
256+
raise ValueError("CandidateFilter or filter_fn must be provided")
257+
return self._candidate_filter.filter_candidates(
258+
candidate_tokens=control_cand,
259+
tokenizer=self._workers[worker_idx].tokenizer,
260+
current_control=current_control_str,
261+
)

‎pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py‎

Lines changed: 19 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
get_embedding_matrix,
1919
get_embeddings,
2020
)
21+
from pyrit.executor.promptgen.gcg.attack.gcg.candidate_proposer import GCGCandidateProposer
2122
from pyrit.executor.promptgen.gcg.default_implementations import (
2223
CrossEntropyLoss,
2324
LengthPreservingFilter,
@@ -308,61 +309,26 @@ def step(
308309
raise ValueError("GCG optimization requires at least one worker")
309310

310311
main_device = self.models[0].device
311-
control_cands = []
312312
loss_function = self._resolve_loss(target_weight=target_weight, control_weight=control_weight)
313313

314-
for j, worker in enumerate(self.workers):
315-
worker(self.prompts[j], ModelWorkerOperation.GRAD)
316-
317-
# Aggregate gradients
318-
grad = None
319-
for j, worker in enumerate(self.workers):
320-
new_grad = worker.results.get().to(main_device)
321-
new_grad = new_grad / new_grad.norm(dim=-1, keepdim=True)
322-
if grad is None:
323-
grad = torch.zeros_like(new_grad)
324-
if grad.shape != new_grad.shape:
325-
with torch.no_grad():
326-
control_cand = self._sample_control_candidates(
327-
worker_index=j - 1,
328-
gradient=grad,
329-
batch_size=batch_size,
330-
topk=topk,
331-
temp=temp,
332-
allow_non_ascii=allow_non_ascii,
333-
)
334-
control_cands.append(
335-
self._filter_control_candidates(
336-
worker_index=j - 1,
337-
control_cand=control_cand,
338-
filter_cand=filter_cand,
339-
)
340-
)
341-
grad = new_grad
342-
else:
343-
grad += new_grad
344-
345-
if grad is None:
346-
raise RuntimeError("GCG workers did not produce an aggregate gradient")
347-
348-
last_worker_index = len(self.workers) - 1
349-
with torch.no_grad():
350-
control_cand = self._sample_control_candidates(
351-
worker_index=last_worker_index,
352-
gradient=grad,
353-
batch_size=batch_size,
354-
topk=topk,
355-
temp=temp,
356-
allow_non_ascii=allow_non_ascii,
357-
)
358-
control_cands.append(
359-
self._filter_control_candidates(
360-
worker_index=last_worker_index,
361-
control_cand=control_cand,
362-
filter_cand=filter_cand,
363-
)
364-
)
365-
del grad, control_cand
314+
proposer = GCGCandidateProposer(
315+
workers=self.workers,
316+
prompts=self.prompts,
317+
sampling=self._resolve_sampling(),
318+
candidate_filter=self._resolve_candidate_filter(filter_cand=filter_cand),
319+
sample_fn=self._sample_control_candidates,
320+
filter_fn=self._filter_control_candidates,
321+
main_device=main_device,
322+
)
323+
candidate_batch = proposer.propose_candidates(
324+
batch_size=batch_size,
325+
topk=topk,
326+
temp=temp,
327+
allow_non_ascii=allow_non_ascii,
328+
filter_cand=filter_cand,
329+
current_control_str=self.control_str,
330+
)
331+
control_cands = candidate_batch.control_candidates_by_group
366332

367333
# Search
368334
loss = torch.zeros(len(control_cands) * batch_size).to(main_device)

0 commit comments

Comments
 (0)