Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
22 changes: 17 additions & 5 deletions pyrit/executor/promptgen/gcg/attack/base/attack_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -316,11 +316,23 @@ def _update_ids(self) -> None:
encoding = self.tokenizer(prompt)
toks = encoding.input_ids

# Locate goal/control/target substrings in the rendered prompt.
goal_start = prompt.find(self.goal)
control_start = prompt.find(self.control)
target_start = prompt.find(self.target)
if goal_start == -1 or control_start == -1 or target_start == -1:
# Locate goal/control/target substrings in the rendered prompt. Searching for each piece
# independently from the start takes the first occurrence anywhere, so a goal that quotes
# its own target (common with affirmative-prefix targets), or a target that also names the
# assistant role marker, silently produced slices pointing at the wrong turn. Instead, find
# where the assistant content starts: rendering only the user turn with a generation
# prompt gives exactly the text before it, provided the full prompt extends that render.
# The control is then the last occurrence before that boundary (it ends the user content),
# the goal the last one before the control, and the target the first one after it.
user_prompt = self.tokenizer.apply_chat_template(messages[:1], tokenize=False, add_generation_prompt=True)
verified = isinstance(user_prompt, str) and len(user_prompt) < len(prompt) and prompt.startswith(user_prompt)
Comment thread
romanlutz marked this conversation as resolved.
Outdated
user_end = len(user_prompt) if verified else len(prompt)
control_start = prompt.rfind(self.control, 0, user_end)
Comment thread
romanlutz marked this conversation as resolved.
Outdated
goal_start = prompt.rfind(self.goal, 0, control_start) if control_start != -1 else -1
# Without a verified boundary, fall back to the end of the control.
assistant_start = user_end if verified else control_start + len(self.control)
target_start = prompt.find(self.target, assistant_start) if goal_start != -1 else -1
if target_start == -1:
raise ValueError(
"Could not locate goal/control/target in chat-templated prompt. "
f"prompt={prompt!r}, goal={self.goal!r}, "
Expand Down
128 changes: 128 additions & 0 deletions tests/unit/executor/promptgen/gcg/test_gcg_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -465,6 +465,72 @@ def test_raises_with_multiple_workers(self) -> None:
)


def _offset_tokenizer(prompt_text: str) -> Any:
"""Build a mock tokenizer that renders ``prompt_text`` and maps characters to tokens.

Each whitespace-delimited run of characters becomes one token, and ``char_to_token``
reports the token containing a character, which is how a fast tokenizer behaves. This
keeps slice assertions meaningful without downloading a real tokenizer.

Args:
prompt_text (str): The already-rendered chat prompt the tokenizer should return.

Returns:
Any: A mock tokenizer suitable for constructing an AttackPrompt.
"""
spans: list[tuple[int, int]] = []
start: int | None = None
for index, char in enumerate(prompt_text):
if char.isspace():
if start is not None:
spans.append((start, index))
start = None
elif start is None:
start = index
if start is not None:
spans.append((start, len(prompt_text)))

def char_to_token(pos: int) -> int | None:
for token_index, (begin, end) in enumerate(spans):
if begin <= pos < end:
return token_index
return None

encoding = MagicMock()
encoding.input_ids = list(range(len(spans)))
encoding.char_to_token.side_effect = char_to_token

tokenizer = MagicMock()
tokenizer.apply_chat_template.return_value = prompt_text
tokenizer.return_value = encoding
return tokenizer


def _fast_tokenizer(chat_template: str) -> Any:
"""
Build a real fast tokenizer offline that registers the chat role markers as special tokens.

Words map to ``[UNK]``, which is enough here: only the character offsets matter for slicing.

Args:
chat_template (str): The Jinja chat template to render with.

Returns:
Any: A ``PreTrainedTokenizerFast`` with ``chat_template`` set.
"""
from tokenizers import Tokenizer, models, pre_tokenizers
from transformers import PreTrainedTokenizerFast

backend = Tokenizer(models.WordLevel({"[UNK]": 0}, unk_token="[UNK]"))
backend.pre_tokenizer = pre_tokenizers.Whitespace()
tokenizer = PreTrainedTokenizerFast(tokenizer_object=backend, unk_token="[UNK]")
tokenizer.add_special_tokens(
{"additional_special_tokens": ["<|user|>", "<|assistant|>", "<|end|>", "<start_of_turn>", "<end_of_turn>"]}
)
tokenizer.chat_template = chat_template
return tokenizer


class TestUpdateIdsErrorPaths:
"""Tests covering the error / fallback paths in AttackPrompt._update_ids."""

Expand Down Expand Up @@ -596,6 +662,68 @@ def char_to_token(pos: int) -> int | None:
assert prompt._target_slice.stop == len(toks)
assert prompt._target_slice.stop > prompt._target_slice.start

def test_target_is_located_after_the_user_turn_when_the_goal_quotes_it(self) -> None:
"""A goal that quotes its own target must not pull the target slice into the user turn.

Affirmative-prefix targets make this realistic: the same text then appears twice in the
rendered prompt, and taking the first occurrence points the target and loss slices at the
user turn instead of the assistant reply.
"""
goal = "Respond with Sure, here is the plan"
control = "! ! ! !"
target = "Sure, here is the plan"
prompt_text = f"<|user|>\n{goal} {control}<|end|>\n<|assistant|>\n{target}<|end|>"

prompt = AttackPrompt(
goal=goal,
target=target,
tokenizer=_offset_tokenizer(prompt_text),
control_init=control,
)

assert prompt._target_slice.start >= prompt._control_slice.stop
assert prompt._assistant_role_slice.start <= prompt._assistant_role_slice.stop

def test_target_that_names_the_assistant_role_marker_is_found_in_the_reply(self) -> None:
"""A target like "assistant" also matches inside ``<|assistant|>``, which is one special token.

Searching right after the user content lands on the role marker and leaves an empty target
slice, so the search has to start where the assistant content does.
"""
tokenizer = _fast_tokenizer(
"{% for m in messages %}<|{{ m['role'] }}|>{{ m['content'] }}<|end|>{% endfor %}"
"{% if add_generation_prompt %}<|assistant|>{% endif %}"
)

prompt = AttackPrompt(goal="Say it", target="assistant", tokenizer=tokenizer, control_init="! !")

ids = tokenizer("<|user|>Say it ! !<|end|><|assistant|>assistant<|end|>").input_ids
# <|user|> Say it ! ! <|end|> <|assistant|> assistant <|end|>
assert prompt._control_slice == slice(3, 5)
assert prompt._target_slice == slice(7, 8)
assert prompt._loss_slice == slice(6, 7)
assert ids[6] == tokenizer.convert_tokens_to_ids("<|assistant|>")

def test_empty_goal_with_a_trimming_template(self) -> None:
"""Target-only datasets use an empty goal, so the user content is " <control>".

A template that trims the content drops that leading space, so the control has to be found on
its own rather than as part of the raw ``f"{goal} {control}"`` string.
"""
tokenizer = _fast_tokenizer(
"{% for m in messages %}<start_of_turn>{{ 'model' if m['role'] == 'assistant' else m['role'] }}\n"
"{{ m['content'] | trim }}<end_of_turn>\n{% endfor %}"
"{% if add_generation_prompt %}<start_of_turn>model\n{% endif %}"
)

prompt = AttackPrompt(goal="", target="Sure, here", tokenizer=tokenizer, control_init="! ! !")

# <start_of_turn> user ! ! ! <end_of_turn> <start_of_turn> model Sure , here <end_of_turn>
assert prompt._goal_slice == slice(2, 2)
assert prompt._control_slice == slice(2, 5)
assert prompt._target_slice == slice(8, 11)
assert prompt._loss_slice == slice(7, 10)


class TestGetWorkersChatTemplateValidation:
"""Tests for the chat-template precondition in get_workers."""
Expand Down