PERF: Speed up GCG candidate evaluation - #2884
Jayanth Sai Yarlagadda (JayYarlagadda) wants to merge 2 commits into
Conversation
02a7556 to
d752b62
Compare
|
@microsoft-github-policy-service agree |
| } | ||
| if attention_mask is not None: | ||
| prefix_kwargs["attention_mask"] = attention_mask[:1, :prefix_length] | ||
| prefix_output = model(**prefix_kwargs) |
There was a problem hiding this comment.
🟡 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.
There was a problem hiding this comment.
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.
d752b62 to
01b092f
Compare
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:
logits_to_keeponly project the positions needed for the target and control losses.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:
typassedWould appreciate your thoughts on the approach, especially the decision to keep prefix caching opt-in.
Addresses #962.