PrivacyJailbreak/easyjailbreak/attacker/TAP_Mehrotra_2023.py
2025-05-15 14:10:22 +08:00

255 lines
13 KiB
Python

r"""
'Tree of Attacks' Recipe
============================================
This module implements a jailbreak method describe in the paper below.
This part of code is based on the code from the paper.
Paper title: Tree of Attacks: Jailbreaking Black-Box LLMs Automatically
arXiv link: https://arxiv.org/abs/2312.02119
Source repository: https://github.com/RICommunity/TAP
"""
import os
import logging
from tqdm import tqdm
from easyjailbreak.attacker import AttackerBase
from easyjailbreak.datasets import JailbreakDataset
from easyjailbreak.datasets.instance import Instance
from easyjailbreak.utils.log_utils import Logger
from easyjailbreak.models.huggingface_model import HuggingfaceModel
from easyjailbreak.models.openai_model import OpenaiModel
####### 4 major components #######
from easyjailbreak.seed.seed_template import SeedTemplate
from easyjailbreak.mutation.generation.IntrospectGeneration import IntrospectGeneration
from easyjailbreak.constraint.DeleteOffTopic import DeleteOffTopic
from easyjailbreak.metrics.Evaluator.Evaluator_GenerativeGetScore import EvaluatorGenerativeGetScore
from easyjailbreak.selector.SelectBasedOnScores import SelectBasedOnScores
r"""
EasyJailbreak TAP class
============================================
"""
__all__ = ['TAP']
target_model_calls = 0
class TAP(AttackerBase):
r"""
Tree of Attack method, an extension of PAIR method. Use 4 phases:
1. Branching
2. Pruning: (phase 1)
3. Query and Access
4. Pruning: (phase 2)
>>> from easyjailbreak.attacker.TAP_Mehrotra_2023 import TAP
>>> from easyjailbreak.models.huggingface_model import from_pretrained
>>> from easyjailbreak.datasets.jailbreak_datasets import JailbreakDataset
>>> from easyjailbreak.datasets.Instance import Instance
>>> attack_model = from_pretrained(model_path_1)
>>> target_model = from_pretrained(model_path_2)
>>> eval_model = from_pretrained(model_path_3)
>>> dataset = JailbreakDataset('AdvBench')
>>> attacker = TAP(attack_model, target_model, eval_model, dataset)
>>> attacker.attack()
>>> attacker.jailbreak_Dataset.save_to_jsonl("./TAP_results.jsonl")
"""
def __init__(self, attack_model, target_model, eval_model, jailbreak_datasets: JailbreakDataset,
tree_width=10, tree_depth=10,root_num=1, branching_factor=4,keep_last_n=3,
max_n_attack_attempts=5, template_file=None,
attack_max_n_tokens=500,
attack_temperature=1,
attack_top_p=0.9,
target_max_n_tokens=150,
target_temperature=1,
target_top_p=1,
judge_max_n_tokens=10,
judge_temperature=1):
"""
initialize TAP, inherit from AttackerBase
:param ~HuggingfaceModel|~OpenaiModel attack_model: LLM for generating jailbreak prompts during Branching(mutation)
:param ~HuggingfaceModel|~OpenaiModel target_model: LLM being attacked to generate adversarial responses
:param ~HuggingfaceModel|~OpenaiModel eval_model: LLM for evaluating during Pruning:phase1(constraint) and Pruning:phase2(select)
:param ~JailbreakDataset jailbreak_datasets: containing instances which conveys the query and reference responses
:param int tree_width: defining the max width of the conversation nodes during Branching(mutation)
:param int tree_depth: defining the max iteration of a single instance
:param int root_num: defining the number of trees or batch of a single instance
:param int branching_factor: defining the number of children nodes generated by a parent node during Branching(mutation)
:param int keep_last_n: defining the number of rounds of dialogue to keep during Branching(mutation)
:param int max_n_attack_attempts: defining the max number of attempts to generating a valid adversarial prompt of a branch
:param str template_file: file path of the seed_template.json
:param int attack_max_n_tokens: max_n_tokens of the target model
:param float attack_temperature: temperature of the attack model
:param float attack_top_p: top p of the attack_model
:param int target_max_n_tokens: max_n_tokens of the target model
:param float target_temperature: temperature of the target model
:param float target_top_p: top_p of the target model
:param int judge_max_n_tokens: max_n_tokens of the target model
:param float judge_temperature: temperature of the judge model
"""
super().__init__(attack_model=attack_model,
target_model=target_model,
eval_model=eval_model,
jailbreak_datasets=jailbreak_datasets)
self.seeds=SeedTemplate().new_seeds(1,method_list=['TAP'],template_file=template_file)
####### 4 major components ##########
self.mutator=IntrospectGeneration(attack_model,
system_prompt=self.seeds[0],
keep_last_n=keep_last_n,
branching_factor=branching_factor,
max_n_attack_attempts=max_n_attack_attempts)
self.constraint=DeleteOffTopic(self.eval_model, tree_width)
self.selector=SelectBasedOnScores(jailbreak_datasets, tree_width)
self.evaluator=EvaluatorGenerativeGetScore(self.eval_model)
######## logging information ############
self.current_query: int = 0
self.current_jailbreak: int = 0
self.current_reject: int = 0
self.current_iteration: int = 0
######## parameters of TAP tree #########
self.root_num = root_num
self.tree_depth = tree_depth
self.tree_width = tree_width
self.branching_factor = branching_factor
######## datasets and logger ############
self.jailbreak_Dataset = JailbreakDataset([])
self.logger = Logger()
######## model configuration ############
self.target_max_n_tokens = target_max_n_tokens
self.target_temperature = target_temperature
self.target_top_p = target_top_p
self.judge_temperature = judge_temperature
self.judge_max_n_tokens = judge_max_n_tokens
if self.attack_model.generation_config == {}:
if isinstance(self.attack_model, OpenaiModel):
self.attack_model.generation_config = {'max_tokens': attack_max_n_tokens,
'temperature': attack_temperature,
'top_p': attack_top_p}
elif isinstance(self.attack_model, HuggingfaceModel):
self.attack_model.generation_config = {'max_new_tokens': attack_max_n_tokens,
'temperature': attack_temperature,
'do_sample': True,
'top_p': attack_top_p,
'eos_token_id': self.attack_model.tokenizer.eos_token_id}
if isinstance(self.eval_model, OpenaiModel) and self.eval_model.generation_config == {}:
self.eval_model.generation_config = {'max_tokens': self.judge_max_n_tokens,
'temperature': self.judge_temperature}
elif isinstance(self.eval_model, HuggingfaceModel) and self.eval_model.generation_config == {}:
self.eval_model.generation_config = {'do_sample': True,
'max_new_tokens': self.judge_max_n_tokens,
'temperature': self.judge_temperature}
def attack(self, save_path='TAP_attack_result.jsonl'):
r"""
Execute the attack process using provided prompts.
"""
# To calculate how many times are eval_model.generate() called
global target_model_calls
logging.info("Jailbreak started!")
try:
for Instance in tqdm(self.jailbreak_datasets, desc="Processing instances"):
new_Instance = self.single_attack(Instance)[0]
self.jailbreak_Dataset.add(new_Instance)
except KeyboardInterrupt:
logging.info("Jailbreak interrupted by user!")
self.update(self.jailbreak_Dataset)
print(f'jailbreak_prompt:{[instance.jailbreak_prompt for instance in self.jailbreak_Dataset]}')
print(f'target_responses:{[instance.target_responses[0] for instance in self.jailbreak_Dataset]}')
print(f"ASR:{100*self.current_jailbreak/self.current_query}%")
print(f"Total calls of generate:{target_model_calls}")
print(f"Eval calls of generate:{self.evaluator.eval_model.generate.count_calls - target_model_calls}")
self.log()
logging.info("Jailbreak finished!")
self.jailbreak_Dataset.save_to_jsonl(save_path)
logging.info(
'Jailbreak result saved at {}!'.format(os.path.join(os.path.dirname(os.path.abspath(__file__)), save_path))
)
def single_attack(self, instance) -> JailbreakDataset:
r"""
Conduct an attack for an instance.
:param ~Instance instance: The Instance that is attacked.
:return ~JailbreakDataset: returns the attack result dataset.
"""
global target_model_calls
batch=[JailbreakDataset([instance.copy()]) for _ in range(self.root_num)]
find_flag = 0
print(f"QUERY:{'='*20}\n{instance.query}")
for iteration in range(1, self.tree_depth + 1):
print(f"""\n{'=' * 36}\nTree-depth is: {iteration}\n{'=' * 36}\n""", flush=True)
dataset_list = []
for i,stream in enumerate(batch):
print(f"BATCH:{i}")
new_dataset = stream
############# generate jailbreak_prompts by branching ################
new_dataset = self.mutator(new_dataset)
############# prune off-topic jailbreak_prompt ################
new_dataset = self.constraint(new_dataset)
############# attack ################
self.target_model.conversation.messages = []
for instance in new_dataset:
if isinstance(self.target_model, OpenaiModel):
instance.target_responses = [
self.target_model.generate(instance.jailbreak_prompt, max_tokens=self.target_max_n_tokens,
temperature=self.target_temperature, top_p=self.target_top_p)]
elif isinstance(self.target_model, HuggingfaceModel):
instance.target_responses = [
self.target_model.generate(instance.jailbreak_prompt,
max_new_tokens=self.target_max_n_tokens,
temperature=self.target_temperature, do_sample=True,
top_p=self.target_top_p,
eos_token_id=self.target_model.tokenizer.eos_token_id)]
target_model_calls+=1
############# prune not-jailbroken jailbreak_prompt ################
num_responses = len(new_dataset)
self.evaluator(new_dataset)
new_dataset = self.selector.select(new_dataset)
print(f"""\n\t{'=' * 36}\n\tCount of Calls of Evaluator is: {self.evaluator.eval_model.generate.calls - num_responses}\n{'=' * 36}\n""", flush=True)
batch[i] = new_dataset
############# attack successful ################
if any([instance.eval_results[-1] == 10 for instance in new_dataset]):
find_flag = 1
print("Found a jailbreak. Exiting.")
break
if find_flag:
new_instance = max(new_dataset, key=lambda instance: instance.eval_results[-1])
new_instance.eval_results=[1]
break
if iteration == self.tree_depth:
new_instance = max(new_dataset, key=lambda instance: instance.eval_results[-1])
new_instance.eval_results=[0]
return JailbreakDataset([new_instance])
def update(self, Dataset: JailbreakDataset):
r"""
Update the state of the ReNeLLM based on the evaluation results of Datasets.
:param ~JailbreakDataset: processed dataset after an iteration
"""
for prompt_node in Dataset:
self.current_jailbreak += prompt_node.num_jailbreak
self.current_query += prompt_node.num_query
self.current_reject += prompt_node.num_reject
def log(self):
r"""
Report the attack results.
"""
logging.info("======Jailbreak report:======")
logging.info(f"Total queries: {self.current_query}")
logging.info(f"Total jailbreak: {self.current_jailbreak}")
logging.info(f"Total reject: {self.current_reject}")
logging.info("========Report End===========")