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

139 lines
5.3 KiB
Python

"""
ICA Class
============================================
This Class executes the In-Context Attack algorithm described in the paper below.
This part of code is based on the paper.
Paper title: Jailbreak and Guard Aligned Language Models with Only Few In-Context Demonstrations
arXiv link: https://arxiv.org/pdf/2310.06387.pdf
"""
import logging
import tqdm
from easyjailbreak.attacker import AttackerBase
from easyjailbreak.datasets import JailbreakDataset
from easyjailbreak.datasets.instance import Instance
from easyjailbreak.seed import SeedTemplate
from easyjailbreak.metrics.Evaluator import EvaluatorPatternJudge
class ICA(AttackerBase):
r"""
In-Context Attack(ICA) crafts malicious contexts to guide models in generating harmful outputs.
"""
def __init__(
self,
target_model,
jailbreak_datasets: JailbreakDataset,
attack_model = None,
eval_model = None,
max_query: int = 100,
max_jailbreak: int = 100,
max_reject: int = 100,
max_iteration: int = 100,
prompt_num: int = 5,
user_input: bool = False,
pattern_dict = None,
):
r"""
Initialize the ICA attack instance.
:param ~model_wrapper target_model: The target model to be attacked.
:param ~JailbreakDataset jailbreak_datasets: The dataset containing harmful queries.
:param ~int prompt_num: The number of in-context demonstration.
:param ~bool user_input: whether to use in-context demonstration input by user.
:param ~dict pattern_dict: the pattern dictionary used in EvaluatorPatternJudge.
"""
super().__init__(attack_model, target_model, eval_model, jailbreak_datasets)
self.attack_results = JailbreakDataset([])
self.evaluator = EvaluatorPatternJudge(pattern_dict=pattern_dict)
self.current_query: int = 0
self.current_jailbreak: int = 0
self.current_reject: int = 0
self.current_iteration: int = 0
self.max_query: int = max_query
self.max_jailbreak: int = max_jailbreak
self.max_reject: int = max_reject
self.max_iteration: int = max_iteration
# ICA特有参数
self.prompt_num: int = prompt_num
self.user_input: bool = user_input
# 初始化jailbreak prompt
if not user_input:
init_prompt = SeedTemplate().new_seeds(seeds_num=1, prompt_usage='attack', method_list=['ICA'])
prompt = init_prompt[0]
else:
harmful_prompts = []
harmful_responses = []
print("Please input " + str(prompt_num) + " pairs of harmful prompts and corresponding responses\n")
for i in range(prompt_num):
harmful_prompts.append(input("harmful prompt:"))
harmful_responses.append(input("harmful response:"))
prompt = ""
for i in range(prompt_num):
prompt += "User:" + harmful_prompts[i] + '\nAssistant:' + harmful_responses[i] + '\n'
prompt += "User:{query}"
for instance in self.jailbreak_datasets:
instance.jailbreak_prompt = prompt
def single_attack(self, sample: Instance):
r"""
Conduct a single attack on sample with n-shot attack demonstrations.
Split the original jailbreak_prompt by roles and merge them into the current conversation_template as in-context demonstration.
"""
prompt = sample.jailbreak_prompt.format(query=sample.query)
prompt_splits = prompt.split("\n")
messages = []
for i in range(0, 2*self.prompt_num, 2):
messages.append(prompt_splits[i].replace("User:", ""))
messages.append(prompt_splits[i+1].replace("Assistant:", ""))
messages.append(prompt_splits[-1].replace("User:", ""))
response = self.target_model.generate(messages=messages)
sample.target_responses.append(response)
return JailbreakDataset([sample])
def update(self, Dataset):
"""
Update the state of the attack.
"""
self.current_iteration += 1
for Instance in Dataset:
self.current_jailbreak += Instance.num_jailbreak
self.current_query += Instance.num_query
self.current_reject += Instance.num_reject
def attack(self):
"""
Main loop for the attack process, iterate through jailbreak_datasets.
"""
logging.info("Jailbreak started!")
try:
for Instance in tqdm.tqdm(self.jailbreak_datasets, desc="processing instance"):
mutated_instance = self.single_attack(Instance)[0]
self.attack_results.add(mutated_instance)
self.evaluator(self.attack_results)
self.update(self.attack_results)
except KeyboardInterrupt:
logging.info("Jailbreak interrupted by user!")
self.log()
logging.info("Jailbreak finished!")
return self.attack_results
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}")