From 51bcc365d483f1d8f57f1094a3934d3e5d2d1b08 Mon Sep 17 00:00:00 2001 From: Vardhman Gupta Date: Sat, 19 Sep 2026 08:20:54 +0530 Subject: [PATCH 1/2] MAINT GCG: model progressive admission transitions (#2665) --- .../gcg/attack/base/attack_manager.py | 90 +++---- .../gcg/attack/base/progressive_schedule.py | 190 +++++++++++++ .../gcg/test_progressive_schedule.py | 249 ++++++++++++++++++ 3 files changed, 471 insertions(+), 58 deletions(-) create mode 100644 pyrit/executor/promptgen/gcg/attack/base/progressive_schedule.py create mode 100644 tests/unit/executor/promptgen/gcg/test_progressive_schedule.py diff --git a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py index cccdd870a8..02e3a385a2 100644 --- a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py +++ b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py @@ -24,6 +24,11 @@ from transformers.models.gpt_neox.modeling_gpt_neox import GPTNeoXForCausalLM from transformers.models.gptj.modeling_gptj import GPTJForCausalLM +from pyrit.executor.promptgen.gcg.attack.base.progressive_schedule import ( + ProgressiveScheduleController, + ProgressiveScheduleState, + ScheduleTransitionAction, +) from pyrit.executor.promptgen.gcg.experiments.log import ( log_gpu_memory, log_loss, @@ -1507,23 +1512,25 @@ def run( }, ) - schedule = ProgressiveScheduleState( - goals_admitted=1 if self.progressive_goals else len(self.goals), - workers_admitted=1 if self.progressive_models else len(self.workers), - stop_inner_on_success=self.progressive_goals, + controller = ProgressiveScheduleController( + total_goals=len(self.goals), + total_workers=len(self.workers), + progressive_goals=self.progressive_goals, + progressive_models=self.progressive_models, + n_steps=n_steps, + control_weight=control_weight, + incr_control=incr_control, + stop_on_success=stop_on_success, + verbose=verbose, ) - # Whether ``schedule.loss`` currently reflects an inner run's measured - # loss, as opposed to the ``inf`` sentinel written when a new round is - # admitted. Tracked explicitly so a legitimately non-finite inner loss - # (non-finite model loss or numeric overflow) is not mistaken for an - # unupdated sentinel value. - loss_is_measured = False - - while schedule.steps_completed < n_steps: + + while not controller.is_complete: + controller.before_inner_run() + schedule = controller.state attack = self.managers["MPA"]( - self.goals[: schedule.goals_admitted], - self.targets[: schedule.goals_admitted], - self.workers[: schedule.workers_admitted], + self.goals[: controller.active_goal_count], + self.targets[: controller.active_goal_count], + self.workers[: controller.active_worker_count], self.control, self.test_prefixes, self.logfile, @@ -1532,17 +1539,15 @@ def run( self.test_targets, self.test_workers, ) - 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, + n_steps=controller.remaining_steps, batch_size=batch_size, topk=topk, temp=temp, allow_non_ascii=allow_non_ascii, target_weight=target_weight, - control_weight=control_weight, + control_weight=controller.control_weight, anneal=anneal, anneal_from=schedule.steps_completed, prev_loss=schedule.loss, @@ -1553,28 +1558,13 @@ def run( random_seed=random_seed, ) control, inner_loss, inner_steps = inner_result - schedule.loss = inner_loss - loss_is_measured = True - - schedule.steps_completed += inner_steps self.control = control - # Once the step budget is spent, stop preparing further rounds: - # admissions and their sentinel resets would strand ``inf`` on - # ``schedule.loss`` for a run that legitimately ends right here. - prepare_next_round = schedule.steps_completed < n_steps - - if schedule.goals_admitted < len(self.goals): - if prepare_next_round: - schedule.goals_admitted += 1 - schedule.loss = np.inf - loss_is_measured = False - elif schedule.workers_admitted < len(self.workers): - if prepare_next_round: - schedule.workers_admitted += 1 - schedule.loss = np.inf - loss_is_measured = False - elif schedule.workers_admitted == len(self.workers) and stop_on_success: + action = controller.advance_after_inner_run( + inner_loss=inner_loss, + inner_steps=inner_steps, + ) + if action == ScheduleTransitionAction.FINALIZE_AND_STOP: self._finalize_progressive_run( attack=attack, step=schedule.steps_completed, @@ -1583,27 +1573,11 @@ def run( verbose=verbose, ) break - elif prepare_next_round and isinstance(control_weight, (int, float)) and incr_control: - if control_weight <= 0.09: - control_weight += 0.01 - schedule.loss = np.inf - loss_is_measured = False - if verbose: - logger.info(f"Control weight increased to {control_weight:.5}") - else: - schedule.stop_inner_on_success = False - - # The inner run must have produced a measured loss whenever any - # optimization happened; guards against silent carry-over regressions. - # Whether the loss was measured is tracked explicitly (a completed - # inner run may legitimately report a non-finite loss), never inferred - # from the numeric value. - if schedule.steps_completed > 0: - assert loss_is_measured, "schedule.loss was never updated by the inner run" - self.last_schedule_state = schedule + controller.validate_post_run() + self.last_schedule_state = controller.state - return self.control, schedule.steps_completed + return self.control, controller.state.steps_completed class IndividualPromptAttack: diff --git a/pyrit/executor/promptgen/gcg/attack/base/progressive_schedule.py b/pyrit/executor/promptgen/gcg/attack/base/progressive_schedule.py new file mode 100644 index 0000000000..99181d4c69 --- /dev/null +++ b/pyrit/executor/promptgen/gcg/attack/base/progressive_schedule.py @@ -0,0 +1,190 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Progressive schedule controller and state models for Greedy Coordinate Gradient (GCG) attacks.""" + +import logging +from dataclasses import dataclass +from enum import Enum, auto + +logger = logging.getLogger(__name__) + + +@dataclass +class ProgressiveScheduleState: + """ + Typed schedule state for ``ProgressiveMultiPromptAttack``. + + Tracks how many goals and workers have been admitted so far, together with + the shared step counter and the loss carried between progressive rounds. + Exposed as ``ProgressiveMultiPromptAttack.last_schedule_state`` after a call + to ``ProgressiveMultiPromptAttack.run``. + """ + + goals_admitted: int + workers_admitted: int + steps_completed: int = 0 + loss: float = float("inf") + stop_inner_on_success: bool = False + + +class ScheduleTransitionAction(Enum): + """Action to take following an inner attack result in progressive scheduling.""" + + CONTINUE = auto() + FINALIZE_AND_STOP = auto() + + +class ProgressiveScheduleController: + """Encapsulates progressive admission, step budget scheduling, and state transitions.""" + + def __init__( + self, + *, + total_goals: int, + total_workers: int, + progressive_goals: bool = True, + progressive_models: bool = True, + n_steps: int = 1000, + control_weight: float | None = None, + incr_control: bool = True, + stop_on_success: bool = True, + verbose: bool = True, + ) -> None: + """ + Initialize the progressive schedule controller. + + Args: + total_goals: Total number of attack goals available for admission. + total_workers: Total number of model workers available for admission. + progressive_goals: Whether goals are admitted progressively one at a time. + progressive_models: Whether models/workers are admitted progressively one at a time. + n_steps: Total step budget across all progressive rounds. + control_weight: Current control weight, or None. + incr_control: Whether to increment control weight when all goals/models are admitted. + stop_on_success: Whether to finalize and stop when fully admitted and success is achieved. + verbose: Whether to log control weight updates. + + Raises: + ValueError: If total_goals or total_workers is not positive. + """ + if total_goals <= 0: + raise ValueError(f"total_goals must be positive, got {total_goals}") + if total_workers <= 0: + raise ValueError(f"total_workers must be positive, got {total_workers}") + + self._total_goals = total_goals + self._total_workers = total_workers + self._n_steps = n_steps + self._control_weight = control_weight + self._incr_control = incr_control + self._stop_on_success = stop_on_success + self._verbose = verbose + + self._state = ProgressiveScheduleState( + goals_admitted=1 if progressive_goals else total_goals, + workers_admitted=1 if progressive_models else total_workers, + stop_inner_on_success=progressive_goals, + ) + self._loss_is_measured = False + + @property + def state(self) -> ProgressiveScheduleState: + """The current progressive schedule state.""" + return self._state + + @property + def is_complete(self) -> bool: + """Whether the overall step budget has been exhausted.""" + return self._state.steps_completed >= self._n_steps + + @property + def remaining_steps(self) -> int: + """The remaining number of optimization steps in the budget.""" + return max(0, self._n_steps - self._state.steps_completed) + + @property + def active_goal_count(self) -> int: + """The number of currently admitted goals.""" + return self._state.goals_admitted + + @property + def active_worker_count(self) -> int: + """The number of currently admitted workers.""" + return self._state.workers_admitted + + @property + def control_weight(self) -> float | None: + """The current control weight.""" + return self._control_weight + + @property + def is_fully_admitted(self) -> bool: + """Whether all goals and workers have been admitted.""" + return self._state.goals_admitted == self._total_goals and self._state.workers_admitted == self._total_workers + + def before_inner_run(self) -> None: + """Prepare schedule state immediately before launching an inner attack.""" + if self.is_fully_admitted: + self._state.stop_inner_on_success = False + + def advance_after_inner_run( + self, + *, + inner_loss: float, + inner_steps: int, + ) -> ScheduleTransitionAction: + """ + Update schedule state and determine the next action after an inner attack round completes. + + Args: + inner_loss: Final loss reported by the inner attack. + inner_steps: Number of steps completed by the inner attack. + + Returns: + ScheduleTransitionAction indicating whether to continue or finalize and stop. + """ + self._state.loss = inner_loss + self._loss_is_measured = True + self._state.steps_completed += inner_steps + + prepare_next_round = self._state.steps_completed < self._n_steps + + if self._state.goals_admitted < self._total_goals: + if prepare_next_round: + self._state.goals_admitted += 1 + self._state.loss = float("inf") + self._loss_is_measured = False + return ScheduleTransitionAction.CONTINUE + + if self._state.workers_admitted < self._total_workers: + if prepare_next_round: + self._state.workers_admitted += 1 + self._state.loss = float("inf") + self._loss_is_measured = False + return ScheduleTransitionAction.CONTINUE + + if self._state.workers_admitted == self._total_workers and self._stop_on_success: + return ScheduleTransitionAction.FINALIZE_AND_STOP + + if prepare_next_round and isinstance(self._control_weight, (int, float)) and self._incr_control: + if self._control_weight <= 0.09: + self._control_weight += 0.01 + self._state.loss = float("inf") + self._loss_is_measured = False + if self._verbose: + logger.info(f"Control weight increased to {self._control_weight:.5}") + else: + self._state.stop_inner_on_success = False + + return ScheduleTransitionAction.CONTINUE + + def validate_post_run(self) -> None: + """ + Validate post-run invariants. + + Raises: + AssertionError: If steps were completed but schedule.loss was never updated by an inner run. + """ + if self._state.steps_completed > 0: + assert self._loss_is_measured, "schedule.loss was never updated by the inner run" diff --git a/tests/unit/executor/promptgen/gcg/test_progressive_schedule.py b/tests/unit/executor/promptgen/gcg/test_progressive_schedule.py new file mode 100644 index 0000000000..36854bd1c4 --- /dev/null +++ b/tests/unit/executor/promptgen/gcg/test_progressive_schedule.py @@ -0,0 +1,249 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +# ruff: noqa: E402 + +import math + +import pytest + +pytest.importorskip( + "pyrit.executor.promptgen.gcg.attack.base.progressive_schedule", + reason="GCG optional dependencies not installed", +) + +from pyrit.executor.promptgen.gcg.attack.base.progressive_schedule import ( + ProgressiveScheduleController, + ScheduleTransitionAction, +) + + +class TestProgressiveScheduleControllerValidation: + """Tests input validation for ProgressiveScheduleController.""" + + def test_raises_on_non_positive_goals(self) -> None: + with pytest.raises(ValueError, match="total_goals must be positive"): + ProgressiveScheduleController(total_goals=0, total_workers=2) + + with pytest.raises(ValueError, match="total_goals must be positive"): + ProgressiveScheduleController(total_goals=-1, total_workers=2) + + def test_raises_on_non_positive_workers(self) -> None: + with pytest.raises(ValueError, match="total_workers must be positive"): + ProgressiveScheduleController(total_goals=2, total_workers=0) + + with pytest.raises(ValueError, match="total_workers must be positive"): + ProgressiveScheduleController(total_goals=2, total_workers=-5) + + +class TestProgressiveScheduleControllerInit: + """Tests initialization and state configuration of ProgressiveScheduleController.""" + + def test_initial_state_progressive_both(self) -> None: + controller = ProgressiveScheduleController( + total_goals=3, + total_workers=2, + progressive_goals=True, + progressive_models=True, + n_steps=100, + ) + assert controller.active_goal_count == 1 + assert controller.active_worker_count == 1 + assert controller.state.stop_inner_on_success is True + assert controller.remaining_steps == 100 + assert controller.is_complete is False + assert controller.is_fully_admitted is False + + def test_initial_state_progressive_goals_only(self) -> None: + controller = ProgressiveScheduleController( + total_goals=3, + total_workers=2, + progressive_goals=True, + progressive_models=False, + n_steps=50, + ) + assert controller.active_goal_count == 1 + assert controller.active_worker_count == 2 + assert controller.state.stop_inner_on_success is True + assert controller.is_fully_admitted is False + + def test_initial_state_progressive_models_only(self) -> None: + controller = ProgressiveScheduleController( + total_goals=3, + total_workers=2, + progressive_goals=False, + progressive_models=True, + n_steps=50, + ) + assert controller.active_goal_count == 3 + assert controller.active_worker_count == 1 + assert controller.state.stop_inner_on_success is False + assert controller.is_fully_admitted is False + + def test_initial_state_no_progressive(self) -> None: + controller = ProgressiveScheduleController( + total_goals=3, + total_workers=2, + progressive_goals=False, + progressive_models=False, + n_steps=50, + ) + assert controller.active_goal_count == 3 + assert controller.active_worker_count == 2 + assert controller.state.stop_inner_on_success is False + assert controller.is_fully_admitted is True + + +class TestProgressiveScheduleControllerTransitions: + """Tests state transitions and admission sequencing in ProgressiveScheduleController.""" + + def test_goal_admission_before_worker_admission(self) -> None: + controller = ProgressiveScheduleController( + total_goals=2, + total_workers=2, + progressive_goals=True, + progressive_models=True, + n_steps=10, + ) + assert controller.active_goal_count == 1 + assert controller.active_worker_count == 1 + + # First inner run completes (budget remaining) + action = controller.advance_after_inner_run(inner_loss=0.5, inner_steps=2) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.active_goal_count == 2 + assert controller.active_worker_count == 1 + assert math.isinf(controller.state.loss) + + # Second inner run completes (all goals admitted, worker should be admitted next) + action = controller.advance_after_inner_run(inner_loss=0.4, inner_steps=2) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.active_goal_count == 2 + assert controller.active_worker_count == 2 + assert controller.is_fully_admitted is True + + def test_finalize_and_stop_action_when_fully_admitted_and_stop_on_success(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + progressive_goals=True, + progressive_models=True, + n_steps=10, + stop_on_success=True, + ) + action = controller.advance_after_inner_run(inner_loss=0.1, inner_steps=3) + assert action == ScheduleTransitionAction.FINALIZE_AND_STOP + assert controller.state.steps_completed == 3 + assert controller.state.loss == 0.1 + + def test_control_weight_ratchet_increments_and_resets_loss(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + progressive_goals=False, + progressive_models=False, + n_steps=10, + control_weight=0.05, + incr_control=True, + stop_on_success=False, + ) + action = controller.advance_after_inner_run(inner_loss=0.8, inner_steps=2) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.control_weight == pytest.approx(0.06) + assert math.isinf(controller.state.loss) + + def test_control_weight_above_threshold_disables_stop_inner_on_success(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + progressive_goals=False, + progressive_models=False, + n_steps=10, + control_weight=0.10, + incr_control=True, + stop_on_success=False, + ) + controller.state.stop_inner_on_success = True + action = controller.advance_after_inner_run(inner_loss=0.8, inner_steps=2) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.control_weight == 0.10 + assert controller.state.stop_inner_on_success is False + + def test_exact_budget_exhaustion_on_goal_boundary_skips_admission(self) -> None: + controller = ProgressiveScheduleController( + total_goals=2, + total_workers=1, + progressive_goals=True, + progressive_models=True, + n_steps=3, + ) + action = controller.advance_after_inner_run(inner_loss=0.75, inner_steps=3) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.is_complete is True + assert controller.active_goal_count == 1 + assert controller.state.loss == 0.75 + + def test_exact_budget_exhaustion_on_worker_boundary_skips_admission(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=2, + progressive_goals=False, + progressive_models=True, + n_steps=3, + ) + action = controller.advance_after_inner_run(inner_loss=0.6, inner_steps=3) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.is_complete is True + assert controller.active_worker_count == 1 + assert controller.state.loss == 0.6 + + def test_exact_budget_exhaustion_on_control_weight_skips_ratchet(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + progressive_goals=False, + progressive_models=False, + n_steps=3, + control_weight=0.05, + incr_control=True, + stop_on_success=False, + ) + action = controller.advance_after_inner_run(inner_loss=0.8, inner_steps=3) + assert action == ScheduleTransitionAction.CONTINUE + assert controller.is_complete is True + assert controller.control_weight == 0.05 + assert controller.state.loss == 0.8 + + def test_non_finite_inner_loss_is_recorded_and_validated(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + n_steps=5, + stop_on_success=False, + ) + controller.advance_after_inner_run(inner_loss=float("inf"), inner_steps=5) + assert controller.state.loss == float("inf") + # Should not raise AssertionError: + controller.validate_post_run() + + def test_validate_post_run_raises_when_steps_completed_without_measured_loss(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + n_steps=5, + ) + controller.state.steps_completed = 3 + with pytest.raises(AssertionError, match="schedule.loss was never updated"): + controller.validate_post_run() + + def test_before_inner_run_disables_stop_inner_on_success_when_fully_admitted(self) -> None: + controller = ProgressiveScheduleController( + total_goals=1, + total_workers=1, + progressive_goals=True, + progressive_models=True, + n_steps=10, + ) + controller.state.stop_inner_on_success = True + controller.before_inner_run() + assert controller.state.stop_inner_on_success is False From 2138122a755d55cf95dda8fa02ad07f6937e5a7a Mon Sep 17 00:00:00 2001 From: Vardhman Gupta Date: Tue, 22 Sep 2026 08:34:18 +0530 Subject: [PATCH 2/2] MAINT GCG: export ProgressiveScheduleState from progressive_schedule in attack_manager (#2665) --- .../gcg/attack/base/attack_manager.py | 18 ------------------ .../executor/promptgen/gcg/test_run_state.py | 10 +++++++++- 2 files changed, 9 insertions(+), 19 deletions(-) diff --git a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py index 02e3a385a2..9914ffd9fc 100644 --- a/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py +++ b/pyrit/executor/promptgen/gcg/attack/base/attack_manager.py @@ -85,24 +85,6 @@ class OptimizationRunState: stop_reason: StopReason | None = None -@dataclass -class ProgressiveScheduleState: - """ - Typed schedule state for ``ProgressiveMultiPromptAttack``. - - Tracks how many goals and workers have been admitted so far, together with - the shared step counter and the loss carried between progressive rounds. - Exposed as ``ProgressiveMultiPromptAttack.last_schedule_state`` after a call - to ``ProgressiveMultiPromptAttack.run``. - """ - - goals_admitted: int - workers_admitted: int - steps_completed: int = 0 - loss: float = float("inf") - stop_inner_on_success: bool = False - - @dataclass class RngBundle: """Per-run RNG state bundle for deterministic GCG execution.""" diff --git a/tests/unit/executor/promptgen/gcg/test_run_state.py b/tests/unit/executor/promptgen/gcg/test_run_state.py index 087cc11496..aa7d02d827 100644 --- a/tests/unit/executor/promptgen/gcg/test_run_state.py +++ b/tests/unit/executor/promptgen/gcg/test_run_state.py @@ -89,6 +89,12 @@ def test_defaults(self) -> None: assert schedule.loss == float("inf") assert schedule.stop_inner_on_success is False + def test_exported_class_identity_matches_progressive_schedule_module(self) -> None: + from pyrit.executor.promptgen.gcg.attack.base import progressive_schedule + + assert ProgressiveScheduleState is progressive_schedule.ProgressiveScheduleState + assert attack_manager_mod.ProgressiveScheduleState is progressive_schedule.ProgressiveScheduleState + class TestMultiPromptRunStateTracking: def test_run_sets_max_steps_reached_when_loop_exhausts(self) -> None: @@ -319,7 +325,9 @@ def test_finalize_phase_logs_final_evaluation_and_stops(self) -> None: control, steps = progressive.run(n_steps=10, stop_on_success=True) assert (control, steps) == ("ctrl", 2) - schedule: ProgressiveScheduleState = progressive.last_schedule_state + schedule = progressive.last_schedule_state + assert isinstance(schedule, attack_manager_mod.ProgressiveScheduleState) + assert isinstance(schedule, ProgressiveScheduleState) assert schedule.steps_completed == 2 assert schedule.goals_admitted == 1 assert schedule.workers_admitted == 1