Skip to content

PERF: Speed up GCG candidate evaluation - #2884

Open
Jayanth Sai Yarlagadda (JayYarlagadda) wants to merge 2 commits into
microsoft:mainfrom
JayYarlagadda:jay/issue-962-prefix-cache
Open

Jayanth Sai Yarlagadda (JayYarlagadda) wants to merge 2 commits into
microsoft:mainfrom
JayYarlagadda:jay/issue-962-prefix-cache

Conversation

@JayYarlagadda

@JayYarlagadda Jayanth Sai Yarlagadda (JayYarlagadda) commented Sep 26, 2026 •

Copy link
Copy Markdown

Hey Roman Lutz (@romanlutz) — I profiled the candidate-evaluation path and found that most of its time and memory were spent creating full-vocabulary logits for every candidate and then passing those tensors back to the parent process.

This PR keeps the existing multi-prompt and multi-worker design, but avoids doing that unnecessary work:

  • The built-in cross-entropy loss is calculated inside each model worker, so the worker returns one loss value per candidate instead of the full logits.
  • Models that support logits_to_keep only project the positions needed for the target and control losses.
  • The unchanged prompt prefix can optionally reuse its KV cache across the candidate batch.

Custom loss implementations still use the existing full-logits path. Prefix caching is disabled by default and can be enabled with GCGAlgorithmConfig.use_prefix_cache.

I benchmarked this on an NVIDIA A40 with Qwen/Qwen2.5-7B-Instruct, using a batch size of 512, top-k 256, and 10 GCG steps. I ran the baseline and changed version three times each and excluded the first two steps of every run from the timing summary.

With one prompt, the median step time went from 7.61s to 5.25s, a 31% improvement. With two prompts and transfer enabled, it went from 13.55s to 8.44s, a 37.7% improvement. Peak GPU memory dropped from roughly 40.5 GiB to 24.9 GiB in both cases.

The baseline and optimized versions selected the same final suffix and completed the same number of optimization steps in all six paired runs. The largest difference in the fp16 loss histories was 0.0022, which is why prefix caching remains opt-in.

Validation:

  • 393 GCG unit tests passed
  • Ruff passed
  • ty passed

Would appreciate your thoughts on the approach, especially the decision to keep prefix caching opt-in.

Addresses #962.

@JayYarlagadda

Copy link
Copy Markdown
Author

@microsoft-github-policy-service agree

Comment thread pyrit/executor/promptgen/gcg/attack/base/attack_manager.py Outdated
}
if attention_mask is not None:
prefix_kwargs["attention_mask"] = attention_mask[:1, :prefix_length]
prefix_output = model(**prefix_kwargs)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Should fix: Skip full-vocabulary projection during the prefix pass. This call leaves logits_to_keep at its model default (0), so it projects and retains logits for every prefix token even though only past_key_values is used. On a long goal, a 2,000-token prefix and 100k-token vocabulary create roughly 400 MB of fp16 logits in each worker, and prefix_output remains live during the suffix forward. Since this path only runs for models that support logits_to_keep, could we request only the last prefix logit (for example logits_to_keep=1) and add a test asserting the prefix forward does not return full-length logits? That preserves the cache while avoiding extra projection and peak memory.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for pointing this out. Fixed in 01b092f.

The prefix forward now uses logits_to_keep=1, and prefix_output is released immediately after cache expansion instead of remaining live during the suffix forward.

The Qwen2 regression test checks that the prefix call returns only one logit position while preserving the cached-loss result.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants