Skip to content
Merged
Show file tree
Hide file tree
Changes from 17 commits
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
8a3d7b4
FIX Propagate GCG random_seed to all RNG sources for deterministic ru…
AmruthVamshi Aug 26, 2026
ee652a9
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Aug 26, 2026
fc0a568
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Aug 27, 2026
ce048cd
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Aug 27, 2026
871f709
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Aug 28, 2026
87b583a
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Aug 31, 2026
c040bd9
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 2, 2026
740423c
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 4, 2026
1c09dc3
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 4, 2026
a951fdc
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 5, 2026
3e9d864
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 11, 2026
dd8f30c
FIX: Wire GCG random_seed through all stochastic operations (#2490)
AmruthVamshi Sep 11, 2026
54f94a8
FIX: Address review round 2, bundle self-creation, device-aware gener…
AmruthVamshi Sep 12, 2026
5ef7aec
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 12, 2026
b6b9b61
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 13, 2026
0ab0f0b
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 14, 2026
436ae8b
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 14, 2026
979ab8a
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 14, 2026
60d3a99
TEST: exercise GCG seed isolation under real overlap; unique per-run …
AmruthVamshi Sep 14, 2026
489fa4d
Merge branch 'main' into fix/gcg-random-seed-deterministic
AmruthVamshi Sep 15, 2026
4ae4b00
FIX centralize GCG RNG bundle creation
romanlutz Sep 16, 2026
7916449
Merge branch 'main' into fix/gcg-random-seed-deterministic
romanlutz Sep 16, 2026
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
79 changes: 78 additions & 1 deletion pyrit/executor/promptgen/gcg/attack/base/attack_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,17 @@ class ProgressiveScheduleState:
stop_inner_on_success: bool = False


@dataclass
class RngBundle:
"""Per-run RNG state bundle for deterministic GCG execution."""

np_rng: np.random.Generator
py_rng: random.Random
torch_gens: dict[int, torch.Generator]
base_seed: int
derived_seeds: dict[int, int]


class NpEncoder(json.JSONEncoder):
"""Encode NumPy scalar and array values for JSON output."""

Expand Down Expand Up @@ -995,17 +1006,31 @@ def run(
log_first: bool = False,
filter_cand: bool = True,
verbose: bool = True,
random_seed: int = 42,
) -> tuple[str, float, int]:
"""
Run iterative optimization.

Returns:
tuple[str, float, int]: The final control, loss, and step count.
"""
rng_bundle = getattr(self, "_rng_bundle", None)
py_rng = rng_bundle.py_rng if rng_bundle else random.Random(random_seed)
if rng_bundle:
self._torch_gens = rng_bundle.torch_gens
else:
workers = getattr(self, "workers", [])
try:
sampling_device = workers[0].model.device
self._torch_gens = {
i: torch.Generator(device=sampling_device).manual_seed(random_seed + i) for i in range(len(workers))
}
except (TypeError, AttributeError, IndexError):
self._torch_gens = {i: torch.Generator().manual_seed(random_seed + i) for i in range(len(workers))}

def acceptance_probability(e: float, e_prime: float, k: int) -> bool:
temperature = max(1 - float(k + 1) / (n_steps + anneal_from), 1.0e-7)
return e_prime < e or math.exp(-(e_prime - e) / temperature) >= random.random()
return e_prime < e or math.exp(-(e_prime - e) / temperature) >= py_rng.random()

if target_weight is None:

Expand Down Expand Up @@ -1378,6 +1403,7 @@ def run(
stop_on_success: bool = True,
verbose: bool = True,
filter_cand: bool = True,
random_seed: int = 42,
) -> tuple[str, int]:
"""
Execute the progressive multi-prompt attack.
Expand Down Expand Up @@ -1409,6 +1435,8 @@ def run(
Whether to print verbose output (default is True)
filter_cand (bool, optional):
Whether to filter candidates whose lengths changed after re-tokenization (default is True)
random_seed (int, optional):
Seed for deterministic random number generation (default is 42)

Returns:
tuple[str, int]: The final control suffix and completed step count.
Expand All @@ -1418,6 +1446,25 @@ def run(
# not keep looking current.
self.last_schedule_state = None

rng_bundle = getattr(self, "_rng_bundle", None)
if rng_bundle is None:
derived_seeds = {i: random_seed + i for i in range(len(self.workers))}
try:
sampling_device = self.workers[0].model.device
torch_gens = {
i: torch.Generator(device=sampling_device).manual_seed(derived_seeds[i])
for i in range(len(self.workers))
}
except (TypeError, AttributeError):
torch_gens = {i: torch.Generator().manual_seed(derived_seeds[i]) for i in range(len(self.workers))}
rng_bundle = RngBundle(
np_rng=np.random.default_rng(random_seed),
py_rng=random.Random(random_seed),
torch_gens=torch_gens,
base_seed=random_seed,
derived_seeds=derived_seeds,
)

_update_attack_log_params(
logfile=self.logfile,
params={
Expand All @@ -1432,6 +1479,8 @@ def run(
"anneal": anneal,
"incr_control": incr_control,
"stop_on_success": stop_on_success,
"random_seed": random_seed,
"derived_seeds": rng_bundle.derived_seeds,
},
)

Expand Down Expand Up @@ -1462,6 +1511,7 @@ def run(
)
if schedule.goals_admitted == len(self.goals) and schedule.workers_admitted == len(self.workers):
schedule.stop_inner_on_success = False
attack._rng_bundle = rng_bundle
inner_result: tuple[str, float, int] = attack.run(
n_steps=n_steps - schedule.steps_completed,
batch_size=batch_size,
Expand All @@ -1477,6 +1527,7 @@ def run(
test_steps=test_steps,
filter_cand=filter_cand,
verbose=verbose,
random_seed=random_seed,
)
control, inner_loss, inner_steps = inner_result
schedule.loss = inner_loss
Expand Down Expand Up @@ -1634,6 +1685,7 @@ def run(
stop_on_success: bool = True,
verbose: bool = True,
filter_cand: bool = True,
random_seed: int = 42,
) -> tuple[str, int]:
"""
Execute the individual-prompt attack.
Expand Down Expand Up @@ -1665,10 +1717,31 @@ def run(
Whether to print verbose output (default is True)
filter_cand (bool, optional):
Whether to filter candidates (default is True)
random_seed (int, optional):
Seed for deterministic random number generation (default is 42)

Returns:
tuple[str, int]: The final control suffix and configured step count.
"""
rng_bundle = getattr(self, "_rng_bundle", None)
if rng_bundle is None:
derived_seeds = {i: random_seed + i for i in range(len(self.workers))}
try:
sampling_device = self.workers[0].model.device
torch_gens = {
i: torch.Generator(device=sampling_device).manual_seed(derived_seeds[i])
for i in range(len(self.workers))
}
except (TypeError, AttributeError):
torch_gens = {i: torch.Generator().manual_seed(derived_seeds[i]) for i in range(len(self.workers))}
rng_bundle = RngBundle(
np_rng=np.random.default_rng(random_seed),
py_rng=random.Random(random_seed),
torch_gens=torch_gens,
base_seed=random_seed,
derived_seeds=derived_seeds,
)

_update_attack_log_params(
logfile=self.logfile,
params={
Expand All @@ -1683,6 +1756,8 @@ def run(
"anneal": anneal,
"incr_control": incr_control,
"stop_on_success": stop_on_success,
"random_seed": random_seed,
"derived_seeds": rng_bundle.derived_seeds,
},
)

Expand All @@ -1703,6 +1778,7 @@ def run(
self.test_targets,
self.test_workers,
)
attack._rng_bundle = rng_bundle
attack.run(
n_steps=n_steps,
batch_size=batch_size,
Expand All @@ -1719,6 +1795,7 @@ def run(
log_first=True,
filter_cand=filter_cand,
verbose=verbose,
random_seed=random_seed,
)

return self.control, n_steps
Expand Down
32 changes: 22 additions & 10 deletions pyrit/executor/promptgen/gcg/attack/gcg/gcg_attack.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import inspect
import logging
from typing import Any

Expand Down Expand Up @@ -104,6 +105,7 @@ def sample_control(
topk: int = 256,
temp: float = 1.0,
allow_non_ascii: bool = True,
torch_generator: torch.Generator | None = None,
) -> torch.Tensor:
"""
Sample new control token candidates based on gradients.
Expand All @@ -114,6 +116,7 @@ def sample_control(
topk (int): Number of top gradient positions to sample from. Defaults to 256.
temp (float): Temperature for sampling. Currently unused but kept for API compatibility. Defaults to 1.0.
allow_non_ascii (bool): Whether to allow non-ASCII tokens. Defaults to True.
torch_generator (torch.Generator | None): Optional generator for deterministic sampling.

Returns:
torch.Tensor: Batch of new candidate control token sequences.
Expand All @@ -127,7 +130,9 @@ def sample_control(
torch.int64
)
new_token_val = torch.gather(
top_indices[new_token_pos], 1, torch.randint(0, topk, (batch_size, 1), device=grad.device)
top_indices[new_token_pos],
1,
torch.randint(0, topk, (batch_size, 1), device=grad.device, generator=torch_generator),
)
return original_control_toks.scatter_(1, new_token_pos.unsqueeze(-1), new_token_val)

Expand Down Expand Up @@ -199,15 +204,22 @@ def _sample_control_candidates(
) -> torch.Tensor:
sampler = self._resolve_sampling()
prompt_manager = self.prompts[worker_index]
return sampler.sample_candidates(
gradient=gradient,
control_tokens=prompt_manager.control_toks,
batch_size=batch_size,
top_k=topk,
temperature=temp,
allow_non_ascii=allow_non_ascii,
non_ascii_tokens=prompt_manager.disallowed_toks,
)
torch_gens = getattr(self, "_torch_gens", None) or {}
torch_gen = torch_gens.get(worker_index)
kwargs: dict[str, Any] = {
"gradient": gradient,
"control_tokens": prompt_manager.control_toks,
"batch_size": batch_size,
"top_k": topk,
"temperature": temp,
"allow_non_ascii": allow_non_ascii,
"non_ascii_tokens": prompt_manager.disallowed_toks,
}
if torch_gen is not None:
sig = inspect.signature(sampler.sample_candidates)
if "torch_generator" in sig.parameters:
kwargs["torch_generator"] = torch_gen
return sampler.sample_candidates(**kwargs)

def _filter_control_candidates(
self,
Expand Down
5 changes: 4 additions & 1 deletion pyrit/executor/promptgen/gcg/default_implementations.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@ def sample_candidates(
temperature: float,
allow_non_ascii: bool,
non_ascii_tokens: torch.Tensor,
torch_generator: torch.Generator | None = None,
) -> torch.Tensor:
"""
Sample ``batch_size`` candidate suffix token sequences.
Expand All @@ -79,6 +80,8 @@ def sample_candidates(
the top-k.
non_ascii_tokens (torch.Tensor): Token ids to exclude when
``allow_non_ascii`` is False.
torch_generator (torch.Generator | None): Optional generator for
deterministic sampling.

Returns:
torch.Tensor: Candidate suffix token sequences with shape
Expand All @@ -99,7 +102,7 @@ def sample_candidates(
new_token_val = torch.gather(
top_indices[new_token_pos],
1,
torch.randint(0, top_k, (batch_size, 1), device=gradient.device),
torch.randint(0, top_k, (batch_size, 1), device=gradient.device, generator=torch_generator),
)
return original_control_tokens.scatter_(1, new_token_pos.unsqueeze(-1), new_token_val)

Expand Down
4 changes: 4 additions & 0 deletions pyrit/executor/promptgen/gcg/extension_protocols.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,7 @@ def sample_candidates(
temperature: float,
allow_non_ascii: bool,
non_ascii_tokens: torch.Tensor,
torch_generator: torch.Generator | None = None,
) -> torch.Tensor:
"""
Sample ``batch_size`` candidate suffix token sequences.
Expand All @@ -101,6 +102,9 @@ def sample_candidates(
non_ascii_tokens (torch.Tensor): Token ids to exclude when
``allow_non_ascii`` is False, shape ``(num_disallowed,)``
and integer dtype.
torch_generator (torch.Generator | None): Optional random number
generator for deterministic sampling. When provided, all
random tensor operations should use this generator.

Returns:
torch.Tensor: Candidate suffix token sequences with shape
Expand Down
Loading