Skip to content
Open
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
89 changes: 65 additions & 24 deletions pyrit/executor/promptgen/gcg/attack/base/attack_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,13 +3,14 @@

from __future__ import annotations

import gc
import json
import logging
import math
import random
import time
from copy import deepcopy
from dataclasses import dataclass
from enum import Enum
from typing import TYPE_CHECKING, Any, cast

import numpy as np
Expand Down Expand Up @@ -383,12 +384,10 @@ def logits(self, model: Any, test_controls: Any = None, return_ids: bool = False

if return_ids:
del locs, test_ids
gc.collect()
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
del ids
gc.collect()
return logits

def target_loss(self, logits: torch.Tensor, ids: torch.Tensor) -> torch.Tensor:
Expand Down Expand Up @@ -598,7 +597,14 @@ def grad(self, model: Any) -> torch.Tensor:
Returns:
torch.Tensor: Aggregated prompt gradients.
"""
return torch.stack([prompt.grad(model) for prompt in self._prompts]).sum(dim=0)
first_gradient = self._prompts[0].grad(model)
if len(self._prompts) == 1:
return first_gradient
result_dtype = first_gradient.dtype
gradient = first_gradient.float() if result_dtype in (torch.float16, torch.bfloat16) else first_gradient.clone()
for prompt in self._prompts[1:]:
gradient.add_(prompt.grad(model).to(dtype=gradient.dtype))
return gradient.to(dtype=result_dtype)

def logits(self, model: Any, test_controls: Any = None, return_ids: bool = False) -> Any:
"""
Expand Down Expand Up @@ -888,7 +894,6 @@ def control_weight_fn(_: int) -> float:

steps += 1
start = time.time()
torch.cuda.empty_cache()
control, loss = self.step(
batch_size=batch_size,
topk=topk,
Expand Down Expand Up @@ -940,14 +945,14 @@ def test(
Jailbreak, exact-match, and loss results.
"""
for j, worker in enumerate(workers):
worker(prompts[j], "test", worker.model)
worker(prompts[j], ModelWorkerOperation.TEST)
model_tests = np.array([worker.results.get() for worker in workers])
model_tests_jb = model_tests[..., 0].tolist()
model_tests_mb = model_tests[..., 1].tolist()
model_tests_loss: list[list[float]] = []
if include_loss:
for j, worker in enumerate(workers):
worker(prompts[j], "test_loss", worker.model)
worker(prompts[j], ModelWorkerOperation.TEST_LOSS)
model_tests_loss = [worker.results.get() for worker in workers]

return model_tests_jb, model_tests_mb, model_tests_loss
Expand Down Expand Up @@ -1781,6 +1786,26 @@ def run(
return total_jb, total_em, test_total_jb, test_total_em, total_outputs, test_total_outputs


class ModelWorkerOperation(str, Enum):
"""A model operation supported by ``ModelWorker``."""

GRAD = "grad"
LOGITS = "logits"
CONTRAST_LOGITS = "contrast_logits"
TEST = "test"
TEST_LOSS = "test_loss"


@dataclass(frozen=True)
class ModelWorkerTask:
"""A spawn-safe model worker task that excludes the worker-owned model."""

obj: Any
operation: ModelWorkerOperation | Callable[..., Any]
args: tuple[Any, ...]
kwargs: dict[str, Any]


class ModelWorker:
"""Run model operations in a dedicated multiprocessing worker."""

Expand All @@ -1802,33 +1827,30 @@ def __init__(
move_to_device = cast("Callable[[torch.device], PreTrainedModel]", model.to)
self.model = move_to_device(torch.device(device)).eval()
self.tokenizer = tokenizer
self.tasks: mp.JoinableQueue[Any] = mp.JoinableQueue()
self.tasks: mp.JoinableQueue[ModelWorkerTask | None] = mp.JoinableQueue()
self.results: mp.JoinableQueue[Any] = mp.JoinableQueue()
self.process: mp.Process | None = None

@staticmethod
def run(model: Any, tasks: mp.JoinableQueue[Any], results: mp.JoinableQueue[Any]) -> None:
def run(
model: Any,
tasks: mp.JoinableQueue[ModelWorkerTask | None],
results: mp.JoinableQueue[Any],
) -> None:
"""Process queued model operations until a stop sentinel arrives."""
model.requires_grad_(False)
model.zero_grad(set_to_none=True)
while True:
task = tasks.get()
if task is None:
tasks.task_done()
break
ob, fn, args, kwargs = task
if fn == "grad":
if task.operation is ModelWorkerOperation.GRAD:
with torch.enable_grad(): # type: ignore[no-untyped-call, unused-ignore]
results.put(ob.grad(*args, **kwargs))
results.put(ModelWorker._execute_task(model=model, task=task))
else:
with torch.no_grad():
if fn == "logits":
results.put(ob.logits(*args, **kwargs))
elif fn == "contrast_logits":
results.put(ob.contrast_logits(*args, **kwargs))
elif fn == "test":
results.put(ob.test(*args, **kwargs))
elif fn == "test_loss":
results.put(ob.test_loss(*args, **kwargs))
else:
results.put(fn(*args, **kwargs))
results.put(ModelWorker._execute_task(model=model, task=task))
tasks.task_done()

def start(self) -> ModelWorker:
Expand Down Expand Up @@ -1856,16 +1878,35 @@ def stop(self) -> ModelWorker:
torch.cuda.empty_cache()
return self

def __call__(self, ob: Any, fn: str, *args: Any, **kwargs: Any) -> ModelWorker:
def __call__(
self,
ob: Any,
operation: ModelWorkerOperation | Callable[..., Any],
*args: Any,
**kwargs: Any,
) -> ModelWorker:
"""
Queue an operation for execution by this worker.

Returns:
ModelWorker: This worker.
"""
self.tasks.put((deepcopy(ob), fn, args, kwargs))
self.tasks.put(ModelWorkerTask(obj=deepcopy(ob), operation=operation, args=args, kwargs=kwargs))
return self

@staticmethod
def _execute_task(*, model: Any, task: ModelWorkerTask) -> Any:
"""
Execute a task with the persistent model when the operation requires one.

Returns:
Any: The operation result.
"""
if isinstance(task.operation, ModelWorkerOperation):
method = getattr(task.obj, task.operation.value)
return method(model, *task.args, **task.kwargs)
return task.operation(*task.args, **task.kwargs)


def get_workers(params: Any, evaluation: bool = False) -> tuple[list[ModelWorker], list[ModelWorker]]:
"""
Expand Down
21 changes: 8 additions & 13 deletions pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import gc
import logging
from typing import Any

Expand All @@ -12,6 +11,7 @@

from pyrit.executor.promptgen.gcg.attack.base.attack_manager import (
AttackPrompt,
ModelWorkerOperation,
MultiPromptAttack,
PromptManager,
get_embedding_matrix,
Expand Down Expand Up @@ -48,7 +48,7 @@ def token_gradients(
torch.Tensor: The gradients of each token in the input_slice with respect to the loss.

Raises:
RuntimeError: If backpropagation does not produce token gradients.
RuntimeError: If autograd does not produce token gradients.
"""
embed_weights = get_embedding_matrix(model)
one_hot = torch.zeros(
Expand All @@ -70,11 +70,10 @@ def token_gradients(
targets = input_ids[target_slice]
loss = nn.CrossEntropyLoss()(logits[0, loss_slice, :], targets)

loss.backward()

if one_hot.grad is None:
raise RuntimeError("Model backward pass did not produce token gradients")
return one_hot.grad.clone()
coordinate_gradient = torch.autograd.grad(loss, one_hot, allow_unused=True)[0]
if coordinate_gradient is None:
raise RuntimeError("Autograd did not produce token gradients")
return coordinate_gradient


class GCGAttackPrompt(AttackPrompt):
Expand Down Expand Up @@ -273,7 +272,7 @@ def step(
loss_function = self._resolve_loss(target_weight=target_weight, control_weight=control_weight)

for j, worker in enumerate(self.workers):
worker(self.prompts[j], "grad", worker.model)
worker(self.prompts[j], ModelWorkerOperation.GRAD)

# Aggregate gradients
grad = None
Expand Down Expand Up @@ -324,7 +323,6 @@ def step(
)
)
del grad, control_cand
gc.collect()

# Search
loss = torch.zeros(len(control_cands) * batch_size).to(main_device)
Expand All @@ -336,7 +334,7 @@ def step(
prompt_indices = progress if progress is not None else range(len(self.prompts[0]))
for i in prompt_indices:
for k, worker in enumerate(self.workers):
worker(self.prompts[k][i], "logits", worker.model, cand, return_ids=True)
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(
loss_function.compute_loss(
Expand All @@ -348,7 +346,6 @@ def step(
for k, (logit, token_ids) in enumerate(zip(logits, ids, strict=True))
)
del logits, ids
gc.collect()

if progress is not None:
progress.set_description(
Expand All @@ -359,9 +356,7 @@ def step(
model_idx = min_idx // batch_size
batch_idx = min_idx % batch_size
next_control, cand_loss = control_cands[model_idx][batch_idx], loss[min_idx]

del control_cands, loss
gc.collect()

current_length = self._get_control_length(control=next_control)
if current_length is not None:
Expand Down
Loading