Add files via upload
This commit is contained in:
parent
0978bb2f1d
commit
1284bb346b
238 changed files with 13931 additions and 3 deletions
78
easyjailbreak/selector/EXP3SelectPolicy.py
Normal file
78
easyjailbreak/selector/EXP3SelectPolicy.py
Normal 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))
|
||||
91
easyjailbreak/selector/MCTSExploreSelectPolicy.py
Normal file
91
easyjailbreak/selector/MCTSExploreSelectPolicy.py
Normal 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))
|
||||
34
easyjailbreak/selector/RandomSelector.py
Normal file
34
easyjailbreak/selector/RandomSelector.py
Normal 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])
|
||||
93
easyjailbreak/selector/ReferenceLossSelector.py
Normal file
93
easyjailbreak/selector/ReferenceLossSelector.py
Normal 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)
|
||||
41
easyjailbreak/selector/RoundRobinSelectPolicy.py
Normal file
41
easyjailbreak/selector/RoundRobinSelectPolicy.py
Normal 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)
|
||||
63
easyjailbreak/selector/SelectBasedOnScores.py
Normal file
63
easyjailbreak/selector/SelectBasedOnScores.py
Normal 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)
|
||||
60
easyjailbreak/selector/UCBSelectPolicy.py
Normal file
60
easyjailbreak/selector/UCBSelectPolicy.py
Normal 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)
|
||||
2
easyjailbreak/selector/__init__.py
Normal file
2
easyjailbreak/selector/__init__.py
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
from .selector import SelectPolicy
|
||||
from .ReferenceLossSelector import ReferenceLossSelector
|
||||
Binary file not shown.
BIN
easyjailbreak/selector/__pycache__/RandomSelector.cpython-39.pyc
Normal file
BIN
easyjailbreak/selector/__pycache__/RandomSelector.cpython-39.pyc
Normal file
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
easyjailbreak/selector/__pycache__/__init__.cpython-39.pyc
Normal file
BIN
easyjailbreak/selector/__pycache__/__init__.cpython-39.pyc
Normal file
Binary file not shown.
BIN
easyjailbreak/selector/__pycache__/selector.cpython-39.pyc
Normal file
BIN
easyjailbreak/selector/__pycache__/selector.cpython-39.pyc
Normal file
Binary file not shown.
52
easyjailbreak/selector/selector.py
Normal file
52
easyjailbreak/selector/selector.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue