Add files via upload

This commit is contained in:
redwyd 2025-05-15 14:10:22 +08:00
commit 1284bb346b
238 changed files with 13931 additions and 3 deletions

View file

@ -0,0 +1,78 @@
"""
EXP3SelectPolicy class
==========================
"""
import numpy as np
from easyjailbreak.selector import SelectPolicy
from easyjailbreak.datasets import Instance, JailbreakDataset
__all__ = ["EXP3SelectPolicy"]
class EXP3SelectPolicy(SelectPolicy):
"""
A selection policy based on the Exponential-weight algorithm for Exploration and Exploitation (EXP3).
This policy is designed for environments with adversarial contexts, balancing between exploring new instances
and exploiting known rewards in a JailbreakDataset.
"""
def __init__(self,
Dataset: JailbreakDataset,
energy: float = 1.0,
gamma: float = 0.05,
alpha: float = 25):
"""
Initializes the EXP3SelectPolicy with a given JailbreakDataset and parameters for the EXP3 algorithm.
:param ~JailbreakDataset Dataset: The dataset from which instances will be selected.
:param float energy: Initial energy level (not used in current implementation).
:param float gamma: Parameter for controlling the exploration-exploitation trade-off.
:param float alpha: Learning rate for the weight updates.
"""
super().__init__(Dataset)
self.energy = energy
self.gamma = gamma
self.alpha = alpha
self.last_choice_index = None
self.initial()
def initial(self):
"""
Initializes or resets the weights and probabilities for each instance in the dataset.
"""
self.weights = [1. for _ in range(len(self.Datasets))]
self.probs = [0. for _ in range(len(self.Datasets))]
def select(self) -> JailbreakDataset:
"""
Selects an instance from the dataset based on the EXP3 algorithm.
:return ~JailbreakDataset: The selected instance from the dataset.
"""
if len(self.Datasets) > len(self.weights):
self.weights.extend([1. for _ in range(len(self.Datasets) - len(self.weights))])
if len(self.Datasets) > len(self.probs):
self.probs.extend([0. for _ in range(len(self.Datasets) - len(self.probs))])
np_weights = np.array(self.weights)
probs = (1 - self.gamma) * np_weights / np_weights.sum() + self.gamma / len(self.Datasets)
self.last_choice_index = np.random.choice(len(self.Datasets), p=probs)
self.Datasets[self.last_choice_index].visited_num += 1
self.probs[self.last_choice_index] = probs[self.last_choice_index]
return JailbreakDataset([self.Datasets[self.last_choice_index]])
def update(self, prompt_nodes: JailbreakDataset):
"""
Updates the weights of the last chosen instance based on the success of the prompts.
:param ~JailbreakDataset prompt_nodes: The dataset containing prompts used for updating weights.
"""
succ_num = sum([prompt_node.num_jailbreak for prompt_node in prompt_nodes])
r = 1 - succ_num / len(prompt_nodes)
x = -1 * r / self.probs[self.last_choice_index]
self.weights[self.last_choice_index] *= np.exp(self.alpha * x / len(self.Datasets))

View file

@ -0,0 +1,91 @@
"""
MCTSExploreSelectPolicy class
================================
"""
import numpy as np
from easyjailbreak.datasets import JailbreakDataset, Instance
from easyjailbreak.selector import SelectPolicy
class MCTSExploreSelectPolicy(SelectPolicy):
"""
This class implements a selection policy based on the Monte Carlo Tree Search (MCTS) algorithm.
It is designed to explore and exploit a dataset of instances for effective jailbreaking of LLMs.
"""
def __init__(self, dataset, inital_prompt_pool, Questions,ratio=0.5, alpha=0.1, beta=0.2):
"""
Initialize the MCTS policy with dataset and parameters for exploration and exploitation.
:param ~JailbreakDataset dataset: The dataset from which instances are to be selected.
:param ~JailbreakDataset initial_prompt_pool: A collection of initial prompts to start the selection process.
:param ~JailbreakDataset Questions: A set of questions or tasks to be addressed by the selected instances.
:param float ratio: The balance between exploration and exploitation (default 0.5).
:param float alpha: Penalty parameter for level adjustment (default 0.1).
:param float beta: Reward scaling factor (default 0.2).
"""
super().__init__(dataset)
self.inital_prompt_pool = inital_prompt_pool
self.Questions = Questions
self.step = 0
self.mctc_select_path = []
self.last_choice_index = None
self.rewards = []
self.ratio = ratio # balance between exploration and exploitation
self.alpha = alpha # penalty for level
self.beta = beta
def select(self) -> JailbreakDataset:
"""
Selects an instance from the dataset using MCTS algorithm.
:return ~JailbreakDataset: The selected instance from the dataset.
"""
self.step += 1
if len(self.Datasets) > len(self.rewards):
self.rewards.extend(
[0 for _ in range(len(self.Datasets) - len(self.rewards))])
self.mctc_select_path = []
cur = max(
self.inital_prompt_pool._dataset,
key=lambda pn:
self.rewards[pn.index] / (pn.visited_num + 1) +
self.ratio * np.sqrt(2 * np.log(self.step) /
(pn.visited_num + 0.01))
)
self.mctc_select_path.append(cur)
while len(cur.children) > 0:
if np.random.rand() < self.alpha:
break
cur = max(
cur.children,
key=lambda pn:
self.rewards[pn.index] / (pn.visited_num + 1) +
self.ratio * np.sqrt(2 * np.log(self.step) /
(pn.visited_num + 0.01))
)
self.mctc_select_path.append(cur)
for pn in self.mctc_select_path:
pn.visited_num += 1
self.last_choice_index = cur.index
return JailbreakDataset([cur])
def update(self, prompt_nodes: JailbreakDataset):
"""
Updates the weights of nodes in the MCTS tree based on their performance.
:param ~JailbreakDataset prompt_nodes: Dataset of prompt nodes to update.
"""
# update weight
succ_num = sum([prompt_node.num_jailbreak
for prompt_node in prompt_nodes])
last_choice_node = self.Datasets[self.last_choice_index]
for prompt_node in reversed(self.mctc_select_path):
reward = succ_num / (len(self.Questions)
* len(prompt_nodes))
self.rewards[prompt_node.index] += reward * max(self.beta, (1 - 0.1 * last_choice_node.level))

View file

@ -0,0 +1,34 @@
"""
RandomSelectPolicy class
============================
"""
import random
from easyjailbreak.selector import SelectPolicy
from easyjailbreak.datasets import JailbreakDataset
__all__ = ["RandomSelectPolicy"]
class RandomSelectPolicy(SelectPolicy):
"""
A selection policy that randomly selects an instance from a JailbreakDataset.
It extends the SelectPolicy abstract base class, providing a concrete implementation
for the random selection strategy.
"""
def __init__(self, Datasets: JailbreakDataset):
"""
Initializes the RandomSelectPolicy with a given JailbreakDataset.
:param ~JailbreakDataset Datasets: The dataset from which instances will be randomly selected.
"""
super().__init__(Datasets)
def select(self) -> JailbreakDataset:
"""
Selects an instance randomly from the dataset and increments its visited count.
:return ~JailbreakDataset: The randomly selected instance from the dataset.
"""
seed = random.choice(self.Datasets._dataset)
seed.visited_num += 1
return JailbreakDataset([seed])

View file

@ -0,0 +1,93 @@
from .selector import SelectPolicy
from ..datasets import JailbreakDataset
from ..datasets import Instance
from ..utils import model_utils
from ..models import WhiteBoxModelBase
import warnings
import torch
import logging
class ReferenceLossSelector(SelectPolicy):
"""
This class implements a selection policy based on the reference loss. It selects instances from a set of parents
based on the minimum loss calculated on their reference target, discarding others.
"""
def __init__(self, model:WhiteBoxModelBase, batch_size=None, is_universal=False):
"""
Initialize the selector with a model and optional configuration settings.
:param ~WhiteBoxModelBase model: The model used for calculating loss.
:param int|None batch_size: The size of each batch for loss calculation. If None, batch_size will be the same as the size of dataset. (default None).
:param bool is_universal: If True, considers the loss of all instances with the same jailbreak_prompt together (default False).
"""
assert isinstance(model, WhiteBoxModelBase)
self.model = model
self.batch_size = batch_size
self.is_universal = is_universal
def select(self, dataset)->JailbreakDataset:
"""
Selects instances from the dataset based on the calculated reference loss.
:param ~JailbreakDataset dataset: The dataset from which instances are to be selected.
:return ~JailbreakDataset: A new dataset containing selected instances with minimum reference loss.
"""
if not self.is_universal and len(dataset.group_by_parents()) > 1:
# 将is_universal=False的情况当作True情况的特例来实现
return JailbreakDataset.merge([self.select(JailbreakDataset(group)) for group in dataset.group_by_parents()])
if self.batch_size is None:
batches = [dataset]
else:
batches = [dataset[i: i+self.batch_size] for i in range(0, len(dataset), self.batch_size)]
# calculate loss on reference response
with torch.no_grad():
for batch in batches:
B = len(batch)
# logging.debug(f'Loss selection: mini-batchsize = {B}')
# encode
batch_input_ids = []
batch_labels = []
is_bad_idx_list = []
for idx, instance in enumerate(batch):
assert len(instance.reference_responses) >= 1
if len(instance.reference_responses) > 1:
warnings.warn(f'传入`ReferenceLossSelector`的每个instance的reference_responses大小都为1,而不是{len(instance.reference_responses)}。将默认使用第一个。')
input_ids, _, _, target_slice = model_utils.encode_trace(self.model, instance.query, instance.jailbreak_prompt, instance.reference_responses[0])
try:
is_bad_idx_list.extend([1 if len(input_ids[0]) != instance.token_id_length else 0])
except AttributeError:
pass
labels = torch.full_like(input_ids, -100)
labels[:, target_slice] = input_ids[:, target_slice]
batch_input_ids.append(input_ids)
batch_labels.append(labels)
batch_input_ids = model_utils.pad_and_stack(batch_input_ids, self.model.pad_token_id)
batch_labels = model_utils.pad_and_stack(batch_labels, -100)
# compute loss values for each instance in batch
batch_loss = model_utils.batch_loss(self.model, batch_input_ids, batch_labels) # B
for idx, instance in enumerate(batch):
if len(is_bad_idx_list) != 0 and is_bad_idx_list[idx]:
instance._loss = float('inf')
else:
instance._loss = batch_loss[idx].item()
# select
best_group = None
best_loss = None
for group in dataset.group_by(lambda x: x.jailbreak_prompt):
total_loss = sum([instance._loss for instance in group])
if best_loss is None or total_loss < best_loss:
best_loss = total_loss
best_group = group
logging.info(f'Loss selection: best loss = {best_loss}')
logging.info(f'Loss Selection: best jailbreak prompt = `{best_group[0].jailbreak_prompt}`')
return JailbreakDataset(best_group)

View file

@ -0,0 +1,41 @@
"""
RoundRobinSelectPolicy class
"""
from easyjailbreak.datasets import JailbreakDataset
from easyjailbreak.selector.selector import SelectPolicy
__all__ = ["RoundRobinSelectPolicy"]
class RoundRobinSelectPolicy(SelectPolicy):
"""
A selection policy that selects instances from a JailbreakDataset in a round-robin manner.
This policy iterates over the dataset, selecting each instance in turn, and then repeats the process.
"""
def __init__(self, Dataset: JailbreakDataset):
"""
Initializes the RoundRobinSelectPolicy with a given JailbreakDataset.
:param ~JailbreakDataset Dataset: The dataset from which instances will be selected in a round-robin fashion.
"""
super().__init__(Dataset)
self.index: int = 0
def select(self) -> JailbreakDataset:
"""
Selects the next instance in the dataset based on a round-robin approach and increments its visited count.
:return ~JailbreakDataset: The selected instance from the dataset.
"""
seed = self.Datasets[self.index]
seed.visited_num += 1
self.index = (self.index + 1) % len(self.Datasets)
return JailbreakDataset([seed])
def update(self, prompt_nodes: JailbreakDataset = None):
"""
Updates the selection index based on the length of the dataset.
:param ~JailbreakDataset prompt_nodes: Not used in this implementation.
"""
self.index = (self.index - 1 + len(self.Datasets)) % len(self.Datasets)

View file

@ -0,0 +1,63 @@
r"""
'SelectBasedOnScores', select those instances whose scores are high(scores are on the extent of jailbreaking),
detail information can be found in the following 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 numpy as np
from typing import List
from easyjailbreak.selector.selector import SelectPolicy
from easyjailbreak.datasets import Instance,JailbreakDataset
class SelectBasedOnScores(SelectPolicy):
"""
This class implements a selection policy based on the scores of instances in a JailbreakDataset.
It selects a subset of instances with high scores, relevant for jailbreaking tasks.
"""
def __init__(self, Dataset: JailbreakDataset, tree_width):
r"""
Initialize the selector with a dataset and a tree width.
:param ~JailbreakDataset Dataset: The dataset from which instances are to be selected.
:param int tree_width: The maximum number of instances to select.
"""
super().__init__(Dataset)
self.tree_width = tree_width
def select(self, dataset:JailbreakDataset) -> JailbreakDataset:
r"""
Selects a subset of instances from the dataset based on their scores.
:param ~JailbreakDataset dataset: The dataset from which instances are to be selected.
:return List[Instance]: A list of selected instances with high evaluation scores.
"""
if dataset != None:
list_dataset = [instance for instance in dataset]
# Ensures that elements with the same score are randomly permuted
np.random.shuffle(list_dataset)
list_dataset.sort(key=lambda x:x.eval_results[-1],reverse=True)
# truncate/select based on judge_scores/instance.eval_results[-1]
width = min(self.tree_width, len(list_dataset))
truncated_list = [list_dataset[i] for i in range(width) if list_dataset[i].eval_results[-1] > 0]
# Ensure that the truncated list has at least two elements
if len(truncated_list) == 0:
truncated_list = [list_dataset[0],list_dataset[1]]
return JailbreakDataset(truncated_list)
else:
list_dataset = [instance for instance in self.Datasets]
# Ensures that elements with the same score are randomly permuted
np.random.shuffle(list_dataset)
list_dataset.sort(key=lambda x: x.eval_results[-1], reverse=True)
# truncate/select based on judge_scores/instance.eval_results[-1]
width = min(self.tree_width, len(list_dataset))
truncated_list = [list_dataset[i] for i in range(width) if list_dataset[i].eval_results[-1] > 0]
# Ensure that the truncated list has at least two elements
if len(truncated_list) == 0:
truncated_list = [list_dataset[0], list_dataset[1]]
return JailbreakDataset(truncated_list)

View file

@ -0,0 +1,60 @@
"""
UCBSelectPolicy class
==========================
"""
import numpy as np
from easyjailbreak.datasets import JailbreakDataset
from easyjailbreak.selector.selector import SelectPolicy
__all__ = ["UCBSelectPolicy"]
class UCBSelectPolicy(SelectPolicy):
"""
A selection policy based on the Upper Confidence Bound (UCB) algorithm. This policy is designed
to balance exploration and exploitation when selecting instances from a JailbreakDataset.
It uses the UCB formula to select instances that either have high rewards or have not been explored much.
"""
def __init__(self,
explore_coeff: float = 1.0,
Dataset: JailbreakDataset = None):
"""
Initializes the UCBSelectPolicy with a given JailbreakDataset and exploration coefficient.
:param float explore_coeff: Coefficient to control the exploration-exploitation balance.
:param ~JailbreakDataset Dataset: The dataset from which instances will be selected.
"""
super().__init__(Dataset)
self.step = 0
self.last_choice_index:int = 0
self.explore_coeff = explore_coeff
self.rewards = [0 for _ in range(len(Dataset))]
def select(self) -> JailbreakDataset:
"""
Selects an instance from the dataset based on the UCB algorithm.
:return ~JailbreakDataset: The selected JailbreakDataset from the dataset.
"""
if len(self.Datasets) > len(self.rewards):
self.rewards.extend([0 for _ in range(len(self.Datasets) - len(self.rewards))])
self.step += 1
scores = np.zeros(len(self.Datasets))
for i, prompt_node in enumerate(self.Datasets):
smooth_visited_num = prompt_node.visited_num + 1
scores[i] = self.rewards[i] / smooth_visited_num + \
self.explore_coeff * np.sqrt(2 * np.log(self.step) / smooth_visited_num)
self.last_choice_index = int(np.argmax(scores))
self.Datasets[self.last_choice_index].visited_num += 1
return JailbreakDataset([self.Datasets[self.last_choice_index]])
def update(self, Dataset: JailbreakDataset):
"""
Updates the rewards for the last selected instance based on the success of the prompts.
:param ~JailbreakDataset Dataset: The dataset containing prompts used for updating rewards.
"""
succ_num = sum([prompt_node.num_jailbreak for prompt_node in Dataset])
self.rewards[self.last_choice_index] += succ_num / len(Dataset)

View file

@ -0,0 +1,2 @@
from .selector import SelectPolicy
from .ReferenceLossSelector import ReferenceLossSelector

View file

@ -0,0 +1,52 @@
"""
SelectPolicy class
=======================
This file contains the implementation of policies for selecting instances from datasets,
specifically tailored for use in easy jailbreak scenarios. It defines abstract base classes
and concrete implementations for selecting instances based on various criteria.
"""
from abc import ABC, abstractmethod
from easyjailbreak.datasets import Instance, JailbreakDataset
__all__ = ["SelectPolicy"]
class SelectPolicy(ABC):
"""
Abstract base class representing a policy for selecting instances from a JailbreakDataset.
It provides a framework for implementing various selection strategies.
"""
def __init__(self, Datasets: JailbreakDataset):
"""
Initializes the SelectPolicy with a given JailbreakDataset.
:param ~JailbreakDataset Datasets: The dataset from which instances will be selected.
"""
self.Datasets = Datasets
for k, instance in enumerate(self.Datasets):
instance.visited_num = 0
instance.index = k
@abstractmethod
def select(self) -> JailbreakDataset:
"""
Abstract method that must be implemented by subclasses to define the selection strategy.
:return ~Instance: The selected instance from the dataset.
"""
raise NotImplementedError(
"SelectPolicy must implement select method.")
def update(self, jailbreak_dataset: JailbreakDataset):
"""
Updates the internal state of the selection policy, if necessary.
:param ~JailbreakDataset jailbreak_dataset: The dataset to update the policy with.
"""
pass
def initial(self):
"""
Initializes or resets any internal state of the selection policy, if necessary.
"""
pass