Use black to format
This commit is contained in:
parent
155d8cf80a
commit
06eac3eaba
4 changed files with 563 additions and 794 deletions
257
hacking.py
257
hacking.py
|
|
@ -10,16 +10,25 @@ from transformers import pipeline, AutoTokenizer, AutoModelForSequenceClassifica
|
||||||
from huggingface_hub import login
|
from huggingface_hub import login
|
||||||
from llm_attacks.minimal_gcg.opt_utils import get_filtered_cands
|
from llm_attacks.minimal_gcg.opt_utils import get_filtered_cands
|
||||||
from words import words
|
from words import words
|
||||||
from utils import minimize_tokens, sample_control, count_tokens, get_combined_score, get_random_words, token_gradients_combined, find_best_word_to_add, words_db
|
from utils import (
|
||||||
|
minimize_tokens,
|
||||||
|
sample_control,
|
||||||
|
count_tokens,
|
||||||
|
get_combined_score,
|
||||||
|
get_random_words,
|
||||||
|
token_gradients_combined,
|
||||||
|
find_best_word_to_add,
|
||||||
|
words_db,
|
||||||
|
)
|
||||||
|
|
||||||
# check if cuda is available
|
# check if cuda is available
|
||||||
cuda_available: bool = torch.cuda.is_available()
|
cuda_available: bool = torch.cuda.is_available()
|
||||||
device: torch.device = torch.device("cuda" if cuda_available else "cpu")
|
device: torch.device = torch.device("cuda" if cuda_available else "cpu")
|
||||||
|
|
||||||
# use token from environment variable
|
# use token from environment variable
|
||||||
login(token=os.getenv("HF_TOKEN"))
|
login(token=os.getenv("HF_TOKEN"))
|
||||||
|
|
||||||
alpha: float = 0.3 # weight between benign maximization and loss minimization
|
alpha: float = 0.3 # weight between benign maximization and loss minimization
|
||||||
min_benign_confidence: float = 0.95 # Minimum benign probability to stop (95%)
|
min_benign_confidence: float = 0.95 # Minimum benign probability to stop (95%)
|
||||||
words_to_inject: int = 1
|
words_to_inject: int = 1
|
||||||
improvement_threshold: float = 0.999
|
improvement_threshold: float = 0.999
|
||||||
|
|
@ -27,7 +36,9 @@ rolling_window_size: int = 10 # Number of past scores to consider
|
||||||
|
|
||||||
# Define a threshold for injecting educational text when optimization stagnates
|
# Define a threshold for injecting educational text when optimization stagnates
|
||||||
stagnation_threshold: int = 15 # Inject after this many iterations without significant improvement
|
stagnation_threshold: int = 15 # Inject after this many iterations without significant improvement
|
||||||
grace_period: int = 50 # Number of iterations to allow after text injection before counting against patience
|
grace_period: int = (
|
||||||
|
50 # Number of iterations to allow after text injection before counting against patience
|
||||||
|
)
|
||||||
max_iterations: int = 10000 # Maximum number of optimization iterations
|
max_iterations: int = 10000 # Maximum number of optimization iterations
|
||||||
patience: int = 10000 # Number of iterations with no improvement before early stopping
|
patience: int = 10000 # Number of iterations with no improvement before early stopping
|
||||||
max_top_scores: int = 10 # Number of top scores to maintain
|
max_top_scores: int = 10 # Number of top scores to maintain
|
||||||
|
|
@ -39,9 +50,11 @@ token_penalty_weight: float = 0.4
|
||||||
min_acceptable_benign: float = 0.50
|
min_acceptable_benign: float = 0.50
|
||||||
|
|
||||||
model_id: str = "meta-llama/Llama-Prompt-Guard-2-86M"
|
model_id: str = "meta-llama/Llama-Prompt-Guard-2-86M"
|
||||||
#model_id = "meta-llama/Llama-Prompt-Guard-2-22M"
|
# model_id = "meta-llama/Llama-Prompt-Guard-2-22M"
|
||||||
tokenizer: AutoTokenizer = AutoTokenizer.from_pretrained(model_id)
|
tokenizer: AutoTokenizer = AutoTokenizer.from_pretrained(model_id)
|
||||||
model: AutoModelForSequenceClassification = AutoModelForSequenceClassification.from_pretrained(model_id)
|
model: AutoModelForSequenceClassification = AutoModelForSequenceClassification.from_pretrained(
|
||||||
|
model_id
|
||||||
|
)
|
||||||
model = model.to(device) # Move model to MPS device
|
model = model.to(device) # Move model to MPS device
|
||||||
|
|
||||||
benign_class: str = "label_0"
|
benign_class: str = "label_0"
|
||||||
|
|
@ -72,40 +85,54 @@ init_token_priority: float = 0.0
|
||||||
general_token_priority: float = 0.95
|
general_token_priority: float = 0.95
|
||||||
|
|
||||||
# Try to use top-performing words from the database for the initial prefix
|
# Try to use top-performing words from the database for the initial prefix
|
||||||
top_words = words_db.get_top_words(limit=init_prefix_words_count, min_uses=1, token_weight=init_token_priority)
|
top_words = words_db.get_top_words(
|
||||||
|
limit=init_prefix_words_count, min_uses=1, token_weight=init_token_priority
|
||||||
|
)
|
||||||
if top_words:
|
if top_words:
|
||||||
print(f"Using {len(top_words)} top-performing words from database for initial prefix")
|
print(f"Using {len(top_words)} top-performing words from database for initial prefix")
|
||||||
# Get words with combined token and improvement prioritization
|
# Get words with combined token and improvement prioritization
|
||||||
initial_words = get_random_words(
|
initial_words = get_random_words(
|
||||||
n=init_prefix_words_count,
|
n=init_prefix_words_count,
|
||||||
min_uses=1, # Words must have been tested at least once
|
min_uses=1, # Words must have been tested at least once
|
||||||
token_priority=init_token_priority
|
token_priority=init_token_priority,
|
||||||
)
|
)
|
||||||
init_prefix: str = " ".join(initial_words)
|
init_prefix: str = " ".join(initial_words)
|
||||||
print(f"Created initial prefix using database-informed words (token priority: {init_token_priority})")
|
print(
|
||||||
|
f"Created initial prefix using database-informed words (token priority: {init_token_priority})"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Fall back to random words if the database doesn't have enough data
|
# Fall back to random words if the database doesn't have enough data
|
||||||
init_prefix: str = " ".join(words[:init_prefix_words_count])
|
init_prefix: str = " ".join(words[:init_prefix_words_count])
|
||||||
print(f"Using random words for initial prefix (no database history available)")
|
print(f"Using random words for initial prefix (no database history available)")
|
||||||
|
|
||||||
#init_prefix = "".join(random.choices(words, k=init_prefix_words_count))
|
# init_prefix = "".join(random.choices(words, k=init_prefix_words_count))
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
global injection_text, text, init_prefix_words_count
|
global injection_text, text, init_prefix_words_count
|
||||||
# Parse command line arguments
|
# Parse command line arguments
|
||||||
parser = argparse.ArgumentParser(description="Prompt hacking tool")
|
parser = argparse.ArgumentParser(description="Prompt hacking tool")
|
||||||
parser.add_argument("--injection", type=str,
|
parser.add_argument(
|
||||||
default=injection_text,
|
"--injection",
|
||||||
help="Injection text to use in the template")
|
type=str,
|
||||||
parser.add_argument("--mandatory-text", type=str,
|
default=injection_text,
|
||||||
default=text,
|
help="Injection text to use in the template",
|
||||||
help="Mandatory text to use in the template")
|
)
|
||||||
parser.add_argument("--init-prefix-words-count", type=int,
|
parser.add_argument(
|
||||||
default=init_prefix_words_count,
|
"--mandatory-text",
|
||||||
help="Number of words to use in the initial prefix")
|
type=str,
|
||||||
|
default=text,
|
||||||
|
help="Mandatory text to use in the template",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--init-prefix-words-count",
|
||||||
|
type=int,
|
||||||
|
default=init_prefix_words_count,
|
||||||
|
help="Number of words to use in the initial prefix",
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
# Update the global parameters based on command line arguments
|
# Update the global parameters based on command line arguments
|
||||||
injection_text = args.injection
|
injection_text = args.injection
|
||||||
text = args.mandatory_text
|
text = args.mandatory_text
|
||||||
|
|
@ -113,33 +140,35 @@ def main():
|
||||||
|
|
||||||
print(f"Injection text: {injection_text}")
|
print(f"Injection text: {injection_text}")
|
||||||
print(f"Mandatory text: {text}")
|
print(f"Mandatory text: {text}")
|
||||||
|
|
||||||
print(f"\nTrying initial prefix: {init_prefix}")
|
print(f"\nTrying initial prefix: {init_prefix}")
|
||||||
|
|
||||||
# Convert initial adversarial string to tokens
|
# Convert initial adversarial string to tokens
|
||||||
best_score: float = float('-inf')
|
best_score: float = float("-inf")
|
||||||
best_prefix: Optional[str] = None
|
best_prefix: Optional[str] = None
|
||||||
adv_prefix: str = init_prefix
|
adv_prefix: str = init_prefix
|
||||||
adv_prefix_tokens: torch.Tensor = tokenizer(adv_prefix, return_tensors="pt", add_special_tokens=False)["input_ids"][0]
|
adv_prefix_tokens: torch.Tensor = tokenizer(
|
||||||
|
adv_prefix, return_tensors="pt", add_special_tokens=False
|
||||||
|
)["input_ids"][0]
|
||||||
adv_prefix_tokens = adv_prefix_tokens.to(device) # Move tokens to MPS device
|
adv_prefix_tokens = adv_prefix_tokens.to(device) # Move tokens to MPS device
|
||||||
control_slice: slice = slice(0, len(adv_prefix_tokens)) # Slice representing the prefix tokens
|
control_slice: slice = slice(0, len(adv_prefix_tokens)) # Slice representing the prefix tokens
|
||||||
|
|
||||||
best_iteration_score: float = float('-inf')
|
best_iteration_score: float = float("-inf")
|
||||||
iterations_without_improvement: int = 0
|
iterations_without_improvement: int = 0
|
||||||
|
|
||||||
# Track both rolling and top scores
|
# Track both rolling and top scores
|
||||||
rolling_scores: List[float] = [] # List to store recent scores
|
rolling_scores: List[float] = [] # List to store recent scores
|
||||||
top_scores: List[float] = [] # List to store top scores
|
top_scores: List[float] = [] # List to store top scores
|
||||||
|
|
||||||
# Track token counts
|
# Track token counts
|
||||||
current_token_count: int = count_tokens(adv_prefix)
|
current_token_count: int = count_tokens(adv_prefix)
|
||||||
min_token_count: int = current_token_count
|
min_token_count: int = current_token_count
|
||||||
|
|
||||||
for i in range(max_iterations):
|
for i in range(max_iterations):
|
||||||
# Prepare input tensors using template
|
# Prepare input tensors using template
|
||||||
full_text = injection_text + adv_prefix + text
|
full_text = injection_text + adv_prefix + text
|
||||||
inputs: Dict[str, torch.Tensor] = tokenizer(full_text, return_tensors="pt")
|
inputs: Dict[str, torch.Tensor] = tokenizer(full_text, return_tensors="pt")
|
||||||
input_ids: torch.Tensor = inputs['input_ids'][0].to(device) # Move input_ids to MPS device
|
input_ids: torch.Tensor = inputs["input_ids"][0].to(device) # Move input_ids to MPS device
|
||||||
|
|
||||||
# Compute gradients using combined approach
|
# Compute gradients using combined approach
|
||||||
coordinate_grad: torch.Tensor = token_gradients_combined(
|
coordinate_grad: torch.Tensor = token_gradients_combined(
|
||||||
|
|
@ -149,7 +178,7 @@ def main():
|
||||||
benign_class=benign_class_idx,
|
benign_class=benign_class_idx,
|
||||||
malicious_class=malicious_class_idx,
|
malicious_class=malicious_class_idx,
|
||||||
alpha=alpha,
|
alpha=alpha,
|
||||||
device=device
|
device=device,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Ensure coordinate_grad is on the correct device and has the right shape
|
# Ensure coordinate_grad is on the correct device and has the right shape
|
||||||
|
|
@ -164,9 +193,9 @@ def main():
|
||||||
new_adv_prefix_toks: torch.Tensor = sample_control(
|
new_adv_prefix_toks: torch.Tensor = sample_control(
|
||||||
adv_prefix_tokens,
|
adv_prefix_tokens,
|
||||||
coordinate_grad,
|
coordinate_grad,
|
||||||
batch_size=32, # Larger batch for more candidates
|
batch_size=32, # Larger batch for more candidates
|
||||||
topk=16, # More options per token
|
topk=16, # More options per token
|
||||||
temp=1.5, # Higher temperature for more exploration
|
temp=1.5, # Higher temperature for more exploration
|
||||||
)
|
)
|
||||||
|
|
||||||
# Convert new tokens to text
|
# Convert new tokens to text
|
||||||
|
|
@ -174,7 +203,7 @@ def main():
|
||||||
tokenizer,
|
tokenizer,
|
||||||
new_adv_prefix_toks,
|
new_adv_prefix_toks,
|
||||||
filter_cand=False,
|
filter_cand=False,
|
||||||
curr_control=adv_prefix
|
curr_control=adv_prefix,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Batch evaluation for all candidates with combined scoring
|
# Batch evaluation for all candidates with combined scoring
|
||||||
|
|
@ -185,25 +214,36 @@ def main():
|
||||||
count_range = max(1, max_count - min_count)
|
count_range = max(1, max_count - min_count)
|
||||||
inputs = tokenizer(candidate_texts, return_tensors="pt", padding=True, truncation=True)
|
inputs = tokenizer(candidate_texts, return_tensors="pt", padding=True, truncation=True)
|
||||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
inputs = {k: v.to(device) for k, v in inputs.items()}
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
logits = model(**inputs).logits
|
logits = model(**inputs).logits
|
||||||
probs = torch.softmax(logits, dim=-1)
|
probs = torch.softmax(logits, dim=-1)
|
||||||
benign_scores = probs[:, benign_class_idx].cpu().numpy()
|
benign_scores = probs[:, benign_class_idx].cpu().numpy()
|
||||||
# Compute normalized loss for each candidate
|
# Compute normalized loss for each candidate
|
||||||
losses = nn.CrossEntropyLoss(reduction="none")(logits, torch.zeros(logits.shape[0], device=device, dtype=torch.long))
|
losses = nn.CrossEntropyLoss(reduction="none")(
|
||||||
normalized_losses = (1.0 / (1.0 + losses.cpu().numpy()))
|
logits,
|
||||||
|
torch.zeros(logits.shape[0], device=device, dtype=torch.long),
|
||||||
|
)
|
||||||
|
normalized_losses = 1.0 / (1.0 + losses.cpu().numpy())
|
||||||
|
|
||||||
# Compute token penalty for each candidate
|
# Compute token penalty for each candidate
|
||||||
token_penalties = [1.0 - ((tc - min_count) / count_range) if count_range > 0 else 0 for tc in token_counts]
|
token_penalties = [
|
||||||
|
1.0 - ((tc - min_count) / count_range) if count_range > 0 else 0
|
||||||
|
for tc in token_counts
|
||||||
|
]
|
||||||
# Compute combined score for each candidate
|
# Compute combined score for each candidate
|
||||||
combined_scores = [
|
combined_scores = [
|
||||||
(alpha * benign_scores[i] + (1 - alpha) * normalized_losses[i]) * (1 - token_penalty_weight + token_penalty_weight * token_penalties[i])
|
(alpha * benign_scores[i] + (1 - alpha) * normalized_losses[i])
|
||||||
|
* (1 - token_penalty_weight + token_penalty_weight * token_penalties[i])
|
||||||
for i in range(len(new_adv_prefix))
|
for i in range(len(new_adv_prefix))
|
||||||
]
|
]
|
||||||
idx = int(max(range(len(combined_scores)), key=lambda i: combined_scores[i]))
|
idx = int(max(range(len(combined_scores)), key=lambda i: combined_scores[i]))
|
||||||
adv_prefix = new_adv_prefix[idx]
|
adv_prefix = new_adv_prefix[idx]
|
||||||
|
|
||||||
# Update the tokens for the next iteration
|
# Update the tokens for the next iteration
|
||||||
adv_prefix_tokens = tokenizer(adv_prefix, return_tensors="pt", add_special_tokens=False)["input_ids"][0]
|
adv_prefix_tokens = tokenizer(
|
||||||
|
adv_prefix, return_tensors="pt", add_special_tokens=False
|
||||||
|
)["input_ids"][0]
|
||||||
adv_prefix_tokens = adv_prefix_tokens.to(device)
|
adv_prefix_tokens = adv_prefix_tokens.to(device)
|
||||||
|
|
||||||
# Check the current classification
|
# Check the current classification
|
||||||
|
|
@ -216,10 +256,14 @@ def main():
|
||||||
predicted_class_id: int = logits.argmax().item()
|
predicted_class_id: int = logits.argmax().item()
|
||||||
benign_score: float = probs[0][benign_class_idx].item()
|
benign_score: float = probs[0][benign_class_idx].item()
|
||||||
benign_percentage: float = benign_score * 100
|
benign_percentage: float = benign_score * 100
|
||||||
malicious_score: float = probs[0][malicious_class_idx].item() if malicious_class_idx is not None else 0
|
malicious_score: float = (
|
||||||
|
probs[0][malicious_class_idx].item() if malicious_class_idx is not None else 0
|
||||||
|
)
|
||||||
|
|
||||||
# Calculate combined score
|
# Calculate combined score
|
||||||
loss: torch.Tensor = nn.CrossEntropyLoss()(logits, torch.zeros(logits.shape[0], device=device).long())
|
loss: torch.Tensor = nn.CrossEntropyLoss()(
|
||||||
|
logits, torch.zeros(logits.shape[0], device=device).long()
|
||||||
|
)
|
||||||
normalized_loss: float = 1.0 / (1.0 + loss.item())
|
normalized_loss: float = 1.0 / (1.0 + loss.item())
|
||||||
current_score: float = alpha * benign_score + (1 - alpha) * normalized_loss
|
current_score: float = alpha * benign_score + (1 - alpha) * normalized_loss
|
||||||
|
|
||||||
|
|
@ -244,9 +288,11 @@ def main():
|
||||||
if current_token_count < min_token_count:
|
if current_token_count < min_token_count:
|
||||||
min_token_count = current_token_count
|
min_token_count = current_token_count
|
||||||
|
|
||||||
print(f"Iteration {i+1}: Class={model.config.id2label[predicted_class_id]} " +
|
print(
|
||||||
f"(benign: {benign_percentage:.2f}%, loss_norm: {normalized_loss:.4f}, " +
|
f"Iteration {i+1}: Class={model.config.id2label[predicted_class_id]} "
|
||||||
f"tokens: {current_token_count}, prefix: {adv_prefix})")
|
+ f"(benign: {benign_percentage:.2f}%, loss_norm: {normalized_loss:.4f}, "
|
||||||
|
+ f"tokens: {current_token_count}, prefix: {adv_prefix})"
|
||||||
|
)
|
||||||
|
|
||||||
if current_score > best_iteration_score:
|
if current_score > best_iteration_score:
|
||||||
# New best score, reset counter
|
# New best score, reset counter
|
||||||
|
|
@ -254,46 +300,75 @@ def main():
|
||||||
iterations_without_improvement = 0
|
iterations_without_improvement = 0
|
||||||
elif current_score >= combined_avg * improvement_threshold:
|
elif current_score >= combined_avg * improvement_threshold:
|
||||||
# Score is close enough to combined average, don't count against patience
|
# Score is close enough to combined average, don't count against patience
|
||||||
print(f" Score within {(1-improvement_threshold)*100:.1f}% of combined average, continuing optimization")
|
print(
|
||||||
|
f" Score within {(1-improvement_threshold)*100:.1f}% of combined average, continuing optimization"
|
||||||
|
)
|
||||||
# Don't increment iterations_without_improvement
|
# Don't increment iterations_without_improvement
|
||||||
else:
|
else:
|
||||||
# Score is significantly worse than combined average, count against patience
|
# Score is significantly worse than combined average, count against patience
|
||||||
iterations_without_improvement += 1
|
iterations_without_improvement += 1
|
||||||
print(f" No significant improvement for {iterations_without_improvement}/{patience} iterations")
|
print(
|
||||||
|
f" No significant improvement for {iterations_without_improvement}/{patience} iterations"
|
||||||
|
)
|
||||||
|
|
||||||
# If we're stagnating but not yet at early stopping threshold, try injecting educational text
|
# If we're stagnating but not yet at early stopping threshold, try injecting educational text
|
||||||
if iterations_without_improvement % stagnation_threshold == 0 and iterations_without_improvement < patience:
|
if (
|
||||||
print(f"\n Optimization stagnating. Looking for words to improve benign rating...")
|
iterations_without_improvement % stagnation_threshold == 0
|
||||||
|
and iterations_without_improvement < patience
|
||||||
|
):
|
||||||
|
print(
|
||||||
|
f"\n Optimization stagnating. Looking for words to improve benign rating..."
|
||||||
|
)
|
||||||
|
|
||||||
# Try to find the best word to add
|
# Try to find the best word to add
|
||||||
new_prefix: Optional[str]
|
new_prefix: Optional[str]
|
||||||
improvement: float
|
improvement: float
|
||||||
new_prefix, improvement = find_best_word_to_add(
|
new_prefix, improvement = find_best_word_to_add(
|
||||||
model, tokenizer, injection_text, adv_prefix, text,
|
model,
|
||||||
benign_class_idx, device=device, num_candidates=len(words),
|
tokenizer,
|
||||||
|
injection_text,
|
||||||
|
adv_prefix,
|
||||||
|
text,
|
||||||
|
benign_class_idx,
|
||||||
|
device=device,
|
||||||
|
num_candidates=len(words),
|
||||||
token_priority=general_token_priority, # Equal weight to token count and improvement
|
token_priority=general_token_priority, # Equal weight to token count and improvement
|
||||||
)
|
)
|
||||||
|
|
||||||
if new_prefix and improvement > 0:
|
if new_prefix and improvement > 0:
|
||||||
# Use the optimized prefix with the best word added
|
# Use the optimized prefix with the best word added
|
||||||
adv_prefix = new_prefix
|
adv_prefix = new_prefix
|
||||||
print(f" Applied optimized prefix with improvement of {improvement:.4f}")
|
print(
|
||||||
|
f" Applied optimized prefix with improvement of {improvement:.4f}"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Fall back to adding random words if no improvement found
|
# Fall back to adding random words if no improvement found
|
||||||
snippet: str = " ".join(get_random_words(words_to_inject, 1, token_priority=general_token_priority))
|
snippet: str = " ".join(
|
||||||
|
get_random_words(
|
||||||
|
words_to_inject,
|
||||||
|
1,
|
||||||
|
token_priority=general_token_priority,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
# Insert the snippet at the beginning
|
# Insert the snippet at the beginning
|
||||||
adv_prefix = snippet + " " + adv_prefix
|
adv_prefix = snippet + " " + adv_prefix
|
||||||
print(f" No improvement found, inserted random words at beginning: '{snippet}'")
|
print(
|
||||||
|
f" No improvement found, inserted random words at beginning: '{snippet}'"
|
||||||
|
)
|
||||||
|
|
||||||
# Update tokens for next iteration
|
# Update tokens for next iteration
|
||||||
adv_prefix_tokens = tokenizer(adv_prefix, return_tensors="pt", add_special_tokens=False)["input_ids"][0]
|
adv_prefix_tokens = tokenizer(
|
||||||
|
adv_prefix, return_tensors="pt", add_special_tokens=False
|
||||||
|
)["input_ids"][0]
|
||||||
adv_prefix_tokens = adv_prefix_tokens.to(device)
|
adv_prefix_tokens = adv_prefix_tokens.to(device)
|
||||||
control_slice = slice(0, len(adv_prefix_tokens))
|
control_slice = slice(0, len(adv_prefix_tokens))
|
||||||
|
|
||||||
# Give the model time to improve with the new text by resetting best score tracking
|
# Give the model time to improve with the new text by resetting best score tracking
|
||||||
best_iteration_score = float('-inf')
|
best_iteration_score = float("-inf")
|
||||||
iterations_without_improvement = max(0, iterations_without_improvement - grace_period)
|
iterations_without_improvement = max(
|
||||||
|
0, iterations_without_improvement - grace_period
|
||||||
|
)
|
||||||
print(f" Reset optimization tracking to give new text time to work")
|
print(f" Reset optimization tracking to give new text time to work")
|
||||||
|
|
||||||
# Early stopping check - only stop if consistently no improvement
|
# Early stopping check - only stop if consistently no improvement
|
||||||
|
|
@ -302,46 +377,69 @@ def main():
|
||||||
break
|
break
|
||||||
|
|
||||||
# Always track the best prefix we've seen, even if not high confidence
|
# Always track the best prefix we've seen, even if not high confidence
|
||||||
if model.config.id2label[predicted_class_id].lower() == benign_class and current_score > best_score:
|
if (
|
||||||
|
model.config.id2label[predicted_class_id].lower() == benign_class
|
||||||
|
and current_score > best_score
|
||||||
|
):
|
||||||
best_score = current_score
|
best_score = current_score
|
||||||
best_prefix = adv_prefix
|
best_prefix = adv_prefix
|
||||||
print(f"New best benign prefix found with score: {best_score:.4f}, benign confidence: {benign_percentage:.2f}%")
|
print(
|
||||||
|
f"New best benign prefix found with score: {best_score:.4f}, benign confidence: {benign_percentage:.2f}%"
|
||||||
|
)
|
||||||
|
|
||||||
# Check if we've found a high confidence benign classification (>95%)
|
# Check if we've found a high confidence benign classification (>95%)
|
||||||
if model.config.id2label[predicted_class_id].lower() == benign_class and benign_score > min_benign_confidence:
|
if (
|
||||||
|
model.config.id2label[predicted_class_id].lower() == benign_class
|
||||||
|
and benign_score > min_benign_confidence
|
||||||
|
):
|
||||||
found_high_confidence_benign = True # Set the flag
|
found_high_confidence_benign = True # Set the flag
|
||||||
best_score = current_score
|
best_score = current_score
|
||||||
best_prefix = adv_prefix
|
best_prefix = adv_prefix
|
||||||
print(f"Found high confidence benign classification ({benign_percentage:.2f}%) at iteration {i+1}! Stopping optimization.")
|
print(
|
||||||
|
f"Found high confidence benign classification ({benign_percentage:.2f}%) at iteration {i+1}! Stopping optimization."
|
||||||
|
)
|
||||||
break # Stop optimizing
|
break # Stop optimizing
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# print stack trace
|
# print stack trace
|
||||||
import traceback
|
import traceback
|
||||||
|
|
||||||
traceback.print_exc()
|
traceback.print_exc()
|
||||||
print(f"Error in iteration {i+1}: {str(e)}")
|
print(f"Error in iteration {i+1}: {str(e)}")
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Print final results
|
# Print final results
|
||||||
if found_high_confidence_benign:
|
if found_high_confidence_benign:
|
||||||
print(f"\n===== SUCCESS! Found a high confidence benign classification (>{min_benign_confidence*100}%) =====")
|
print(
|
||||||
|
f"\n===== SUCCESS! Found a high confidence benign classification (>{min_benign_confidence*100}%) ====="
|
||||||
|
)
|
||||||
|
|
||||||
# Now minimize the tokens while maintaining the benign classification
|
# Now minimize the tokens while maintaining the benign classification
|
||||||
if best_prefix is not None:
|
if best_prefix is not None:
|
||||||
best_prefix = minimize_tokens(
|
best_prefix = minimize_tokens(
|
||||||
model, tokenizer, injection_text, best_prefix, text,
|
model,
|
||||||
benign_class_idx, device=device,
|
tokenizer,
|
||||||
|
injection_text,
|
||||||
|
best_prefix,
|
||||||
|
text,
|
||||||
|
benign_class_idx,
|
||||||
|
device=device,
|
||||||
min_acceptable_benign=min_acceptable_benign,
|
min_acceptable_benign=min_acceptable_benign,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
print("\n===== Did not find a high confidence benign classification =====")
|
print("\n===== Did not find a high confidence benign classification =====")
|
||||||
if best_prefix is not None:
|
if best_prefix is not None:
|
||||||
print(f"Best prefix found had a benign score of: {best_score:.4f}")
|
print(f"Best prefix found had a benign score of: {best_score:.4f}")
|
||||||
|
|
||||||
# Still try to minimize tokens
|
# Still try to minimize tokens
|
||||||
best_prefix = minimize_tokens(
|
best_prefix = minimize_tokens(
|
||||||
model, tokenizer, injection_text, best_prefix, text,
|
model,
|
||||||
benign_class_idx, device=device,
|
tokenizer,
|
||||||
|
injection_text,
|
||||||
|
best_prefix,
|
||||||
|
text,
|
||||||
|
benign_class_idx,
|
||||||
|
device=device,
|
||||||
min_acceptable_benign=min_acceptable_benign,
|
min_acceptable_benign=min_acceptable_benign,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -355,8 +453,10 @@ def main():
|
||||||
probs: torch.Tensor = torch.softmax(logits, dim=-1)
|
probs: torch.Tensor = torch.softmax(logits, dim=-1)
|
||||||
predicted_class_id: int = logits.argmax().item()
|
predicted_class_id: int = logits.argmax().item()
|
||||||
benign_prob: float = probs[0][benign_class_idx].item() * 100
|
benign_prob: float = probs[0][benign_class_idx].item() * 100
|
||||||
print(f'\nPayload is: {text}')
|
print(f"\nPayload is: {text}")
|
||||||
print(f'Payload without prefix is classified as: {model.config.id2label[predicted_class_id]} (benign probability: {benign_prob:.2f}%)')
|
print(
|
||||||
|
f"Payload without prefix is classified as: {model.config.id2label[predicted_class_id]} (benign probability: {benign_prob:.2f}%)"
|
||||||
|
)
|
||||||
|
|
||||||
# With adversarial prefix
|
# With adversarial prefix
|
||||||
full_text = injection_text + adv_prefix + text
|
full_text = injection_text + adv_prefix + text
|
||||||
|
|
@ -366,15 +466,17 @@ def main():
|
||||||
probs = torch.softmax(logits, dim=-1)
|
probs = torch.softmax(logits, dim=-1)
|
||||||
predicted_class_id = logits.argmax().item()
|
predicted_class_id = logits.argmax().item()
|
||||||
benign_prob = probs[0][benign_class_idx].item() * 100
|
benign_prob = probs[0][benign_class_idx].item() * 100
|
||||||
print(f'Prefix is: {adv_prefix}')
|
print(f"Prefix is: {adv_prefix}")
|
||||||
print(f'Complete text using template is: {full_text}')
|
print(f"Complete text using template is: {full_text}")
|
||||||
print(f'Payload with prefix is classified as: {model.config.id2label[predicted_class_id]} (benign probability: {benign_prob:.2f}%)')
|
print(
|
||||||
|
f"Payload with prefix is classified as: {model.config.id2label[predicted_class_id]} (benign probability: {benign_prob:.2f}%)"
|
||||||
|
)
|
||||||
|
|
||||||
# Try to run inference with the model's classifier
|
# Try to run inference with the model's classifier
|
||||||
classifier = pipeline(
|
classifier = pipeline(
|
||||||
"text-classification",
|
"text-classification",
|
||||||
model=model_id,
|
model=model_id,
|
||||||
device=0 if torch.cuda.is_available() else -1
|
device=0 if torch.cuda.is_available() else -1,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
@ -388,8 +490,9 @@ def main():
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error running classifier pipeline: {str(e)}")
|
print(f"Error running classifier pipeline: {str(e)}")
|
||||||
|
|
||||||
print(f'Adv prefix token count: {count_tokens(adv_prefix)}')
|
print(f"Adv prefix token count: {count_tokens(adv_prefix)}")
|
||||||
print(f'Total token count: {count_tokens(full_text)}')
|
print(f"Total token count: {count_tokens(full_text)}")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
main()
|
main()
|
||||||
|
|
|
||||||
334
utils.py
334
utils.py
|
|
@ -10,24 +10,25 @@ from wordsdb import WordsDatabase
|
||||||
# Create a global instance of the database
|
# Create a global instance of the database
|
||||||
words_db = WordsDatabase()
|
words_db = WordsDatabase()
|
||||||
|
|
||||||
|
|
||||||
def find_best_word_to_add(
|
def find_best_word_to_add(
|
||||||
model: AutoModelForSequenceClassification,
|
model: AutoModelForSequenceClassification,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
injection_text: str,
|
injection_text: str,
|
||||||
adv_prefix: str,
|
adv_prefix: str,
|
||||||
text: str,
|
text: str,
|
||||||
benign_class_idx: int,
|
benign_class_idx: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
num_candidates: int = 20,
|
num_candidates: int = 20,
|
||||||
token_weight: float = 0.5, # Weight for token count prioritization
|
token_weight: float = 0.5, # Weight for token count prioritization
|
||||||
use_db: bool = True, # Whether to use the database for word selection and tracking
|
use_db: bool = True, # Whether to use the database for word selection and tracking
|
||||||
token_priority: float = 0.3, # How much to prioritize words with fewer tokens when selecting from database
|
token_priority: float = 0.3, # How much to prioritize words with fewer tokens when selecting from database
|
||||||
order_template: str = "{injection}{prefix}{text}" # Template for ordering components
|
order_template: str = "{injection}{prefix}{text}", # Template for ordering components
|
||||||
) -> Tuple[Optional[str], float]:
|
) -> Tuple[Optional[str], float]:
|
||||||
"""
|
"""
|
||||||
Evaluate multiple candidate words and find the one that most improves the benign score when added to the prefix.
|
Evaluate multiple candidate words and find the one that most improves the benign score when added to the prefix.
|
||||||
Prioritizes words that result in fewer tokens while still improving the benign score.
|
Prioritizes words that result in fewer tokens while still improving the benign score.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
-----------
|
-----------
|
||||||
model: The model to evaluate with
|
model: The model to evaluate with
|
||||||
|
|
@ -41,14 +42,14 @@ def find_best_word_to_add(
|
||||||
use_db: Whether to use the database for word selection and tracking
|
use_db: Whether to use the database for word selection and tracking
|
||||||
token_priority: How much to prioritize words with fewer tokens when selecting from database
|
token_priority: How much to prioritize words with fewer tokens when selecting from database
|
||||||
order_template: Template string for ordering components (using {injection}, {prefix}, {text})
|
order_template: Template string for ordering components (using {injection}, {prefix}, {text})
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
--------
|
--------
|
||||||
best_word: The word that most improves the benign score
|
best_word: The word that most improves the benign score
|
||||||
improvement: The amount of improvement in benign score
|
improvement: The amount of improvement in benign score
|
||||||
"""
|
"""
|
||||||
print(f"\n----- TESTING {num_candidates} CANDIDATE WORDS TO ADD (BATCHED) -----")
|
print(f"\n----- TESTING {num_candidates} CANDIDATE WORDS TO ADD (BATCHED) -----")
|
||||||
|
|
||||||
# Get baseline benign score with current prefix
|
# Get baseline benign score with current prefix
|
||||||
try:
|
try:
|
||||||
full_text = order_template.format(injection=injection_text, prefix=adv_prefix, text=text)
|
full_text = order_template.format(injection=injection_text, prefix=adv_prefix, text=text)
|
||||||
|
|
@ -62,22 +63,24 @@ def find_best_word_to_add(
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error testing baseline: {e}")
|
print(f"Error testing baseline: {e}")
|
||||||
return None, 0
|
return None, 0
|
||||||
|
|
||||||
# Generate candidate words to test - prioritize known good words if using database
|
# Generate candidate words to test - prioritize known good words if using database
|
||||||
if use_db:
|
if use_db:
|
||||||
# Try to get high-performing words from the database, with token count consideration
|
# Try to get high-performing words from the database, with token count consideration
|
||||||
db_candidates_count = num_candidates // 2
|
db_candidates_count = num_candidates // 2
|
||||||
if db_candidates_count > 0:
|
if db_candidates_count > 0:
|
||||||
top_words = words_db.get_top_words(
|
top_words = words_db.get_top_words(
|
||||||
limit=db_candidates_count,
|
limit=db_candidates_count,
|
||||||
min_uses=1, # Only need to have been tested once
|
min_uses=1, # Only need to have been tested once
|
||||||
sort_by="combined" if token_priority > 0 else "improvement",
|
sort_by="combined" if token_priority > 0 else "improvement",
|
||||||
token_weight=token_priority
|
token_weight=token_priority,
|
||||||
)
|
)
|
||||||
|
|
||||||
# If we got some words from the database, use them plus some random words
|
# If we got some words from the database, use them plus some random words
|
||||||
if top_words:
|
if top_words:
|
||||||
print(f"Using {len(top_words)} words from database (with token priority {token_priority}) plus {num_candidates - len(top_words)} random words")
|
print(
|
||||||
|
f"Using {len(top_words)} words from database (with token priority {token_priority}) plus {num_candidates - len(top_words)} random words"
|
||||||
|
)
|
||||||
remaining = num_candidates - len(top_words)
|
remaining = num_candidates - len(top_words)
|
||||||
candidates = top_words + random.choices(words, k=remaining)
|
candidates = top_words + random.choices(words, k=remaining)
|
||||||
else:
|
else:
|
||||||
|
|
@ -88,10 +91,10 @@ def find_best_word_to_add(
|
||||||
else:
|
else:
|
||||||
# Just use random words if not using the database
|
# Just use random words if not using the database
|
||||||
candidates = random.choices(words, k=num_candidates)
|
candidates = random.choices(words, k=num_candidates)
|
||||||
|
|
||||||
# Define positions to test for each word
|
# Define positions to test for each word
|
||||||
insert_positions: List[str] = ["beginning", "middle", "end"]
|
insert_positions: List[str] = ["beginning", "middle", "end"]
|
||||||
|
|
||||||
# Generate all candidate prefixes - one for each word + position combination
|
# Generate all candidate prefixes - one for each word + position combination
|
||||||
all_candidate_prefixes = []
|
all_candidate_prefixes = []
|
||||||
for word in candidates:
|
for word in candidates:
|
||||||
|
|
@ -103,33 +106,37 @@ def find_best_word_to_add(
|
||||||
test_prefix = adv_prefix + " " + word
|
test_prefix = adv_prefix + " " + word
|
||||||
else: # middle
|
else: # middle
|
||||||
# Find a reasonable spot to insert in the middle if possible
|
# Find a reasonable spot to insert in the middle if possible
|
||||||
if ' ' in adv_prefix:
|
if " " in adv_prefix:
|
||||||
words_list: List[str] = adv_prefix.split()
|
words_list: List[str] = adv_prefix.split()
|
||||||
middle_idx: int = len(words_list) // 2
|
middle_idx: int = len(words_list) // 2
|
||||||
words_list.insert(middle_idx, word)
|
words_list.insert(middle_idx, word)
|
||||||
test_prefix = ' '.join(words_list)
|
test_prefix = " ".join(words_list)
|
||||||
else:
|
else:
|
||||||
# If no spaces, insert at midpoint of string
|
# If no spaces, insert at midpoint of string
|
||||||
middle_idx: int = len(adv_prefix) // 2
|
middle_idx: int = len(adv_prefix) // 2
|
||||||
test_prefix = adv_prefix[:middle_idx] + " " + word + " " + adv_prefix[middle_idx:]
|
test_prefix = (
|
||||||
|
adv_prefix[:middle_idx] + " " + word + " " + adv_prefix[middle_idx:]
|
||||||
all_candidate_prefixes.append({
|
)
|
||||||
"prefix": test_prefix,
|
|
||||||
"word": word,
|
all_candidate_prefixes.append(
|
||||||
"position": position,
|
{
|
||||||
"token_count": len(tokenizer.encode(test_prefix, add_special_tokens=False))
|
"prefix": test_prefix,
|
||||||
})
|
"word": word,
|
||||||
|
"position": position,
|
||||||
|
"token_count": len(tokenizer.encode(test_prefix, add_special_tokens=False)),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
# Prepare all candidate full texts for batch evaluation
|
# Prepare all candidate full texts for batch evaluation
|
||||||
candidate_full_texts = [
|
candidate_full_texts = [
|
||||||
order_template.format(injection=injection_text, prefix=c["prefix"], text=text)
|
order_template.format(injection=injection_text, prefix=c["prefix"], text=text)
|
||||||
for c in all_candidate_prefixes
|
for c in all_candidate_prefixes
|
||||||
]
|
]
|
||||||
|
|
||||||
if not candidate_full_texts:
|
if not candidate_full_texts:
|
||||||
print("No candidate prefixes to evaluate")
|
print("No candidate prefixes to evaluate")
|
||||||
return None, 0
|
return None, 0
|
||||||
|
|
||||||
# Batch inference
|
# Batch inference
|
||||||
try:
|
try:
|
||||||
inputs = tokenizer(candidate_full_texts, return_tensors="pt", padding=True, truncation=True)
|
inputs = tokenizer(candidate_full_texts, return_tensors="pt", padding=True, truncation=True)
|
||||||
|
|
@ -141,33 +148,33 @@ def find_best_word_to_add(
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error in batch evaluation: {e}")
|
print(f"Error in batch evaluation: {e}")
|
||||||
return None, 0
|
return None, 0
|
||||||
|
|
||||||
# Calculate token counts for normalization
|
# Calculate token counts for normalization
|
||||||
token_counts = [c["token_count"] for c in all_candidate_prefixes]
|
token_counts = [c["token_count"] for c in all_candidate_prefixes]
|
||||||
max_token_count = max(token_counts) if token_counts else 1
|
max_token_count = max(token_counts) if token_counts else 1
|
||||||
|
|
||||||
# Process the results
|
# Process the results
|
||||||
results = []
|
results = []
|
||||||
best_combined_score = 0
|
best_combined_score = 0
|
||||||
best_result_idx = -1
|
best_result_idx = -1
|
||||||
|
|
||||||
print(f"Running score analysis for {len(all_candidate_prefixes)} candidate prefixes")
|
print(f"Running score analysis for {len(all_candidate_prefixes)} candidate prefixes")
|
||||||
|
|
||||||
for idx, candidate in enumerate(all_candidate_prefixes):
|
for idx, candidate in enumerate(all_candidate_prefixes):
|
||||||
benign_score = benign_scores[idx]
|
benign_score = benign_scores[idx]
|
||||||
improvement = benign_score - baseline_score
|
improvement = benign_score - baseline_score
|
||||||
token_count = candidate["token_count"]
|
token_count = candidate["token_count"]
|
||||||
|
|
||||||
# Calculate token efficiency (lower token count is better)
|
# Calculate token efficiency (lower token count is better)
|
||||||
# Normalize token count to 0-1 scale (where 1 is better = fewer tokens)
|
# Normalize token count to 0-1 scale (where 1 is better = fewer tokens)
|
||||||
token_efficiency = 1.0 - min(1.0, token_count / max_token_count)
|
token_efficiency = 1.0 - min(1.0, token_count / max_token_count)
|
||||||
|
|
||||||
# Calculate combined score (weighting improvement and token efficiency)
|
# Calculate combined score (weighting improvement and token efficiency)
|
||||||
# Only consider token efficiency if improvement is positive
|
# Only consider token efficiency if improvement is positive
|
||||||
combined_score = 0
|
combined_score = 0
|
||||||
if improvement > 0:
|
if improvement > 0:
|
||||||
combined_score = (1 - token_weight) * improvement + token_weight * token_efficiency
|
combined_score = (1 - token_weight) * improvement + token_weight * token_efficiency
|
||||||
|
|
||||||
# Record performance in results list
|
# Record performance in results list
|
||||||
result = {
|
result = {
|
||||||
"word": candidate["word"],
|
"word": candidate["word"],
|
||||||
|
|
@ -177,55 +184,64 @@ def find_best_word_to_add(
|
||||||
"tokens": token_count,
|
"tokens": token_count,
|
||||||
"token_efficiency": token_efficiency,
|
"token_efficiency": token_efficiency,
|
||||||
"combined_score": combined_score,
|
"combined_score": combined_score,
|
||||||
"prefix": candidate["prefix"]
|
"prefix": candidate["prefix"],
|
||||||
}
|
}
|
||||||
|
|
||||||
results.append(result)
|
results.append(result)
|
||||||
|
|
||||||
# Record the performance in the database if enabled
|
# Record the performance in the database if enabled
|
||||||
if use_db and improvement != 0: # Only record non-zero improvements
|
if use_db and improvement != 0: # Only record non-zero improvements
|
||||||
words_db.record_word_performance(
|
words_db.record_word_performance(
|
||||||
candidate["word"], candidate["position"], benign_score, improvement,
|
candidate["word"],
|
||||||
token_count, combined_score
|
candidate["position"],
|
||||||
|
benign_score,
|
||||||
|
improvement,
|
||||||
|
token_count,
|
||||||
|
combined_score,
|
||||||
)
|
)
|
||||||
|
|
||||||
#print(f"Word '{candidate['word']}' at {candidate['position']}: {benign_score:.4f} (Δ: {improvement:.4f}, tokens: {token_count}, combined: {combined_score:.4f})")
|
# print(f"Word '{candidate['word']}' at {candidate['position']}: {benign_score:.4f} (Δ: {improvement:.4f}, tokens: {token_count}, combined: {combined_score:.4f})")
|
||||||
|
|
||||||
# Only consider improvements (benign_score > baseline_score)
|
# Only consider improvements (benign_score > baseline_score)
|
||||||
if improvement > 0 and combined_score > best_combined_score:
|
if improvement > 0 and combined_score > best_combined_score:
|
||||||
best_combined_score = combined_score
|
best_combined_score = combined_score
|
||||||
best_result_idx = idx
|
best_result_idx = idx
|
||||||
|
|
||||||
# Sort results by combined score
|
# Sort results by combined score
|
||||||
results.sort(key=lambda x: x["combined_score"], reverse=True)
|
results.sort(key=lambda x: x["combined_score"], reverse=True)
|
||||||
|
|
||||||
# Print top 5 results
|
# Print top 5 results
|
||||||
print("\nTop 5 most effective additions (based on combined score):")
|
print("\nTop 5 most effective additions (based on combined score):")
|
||||||
for i, result in enumerate(results[:5]):
|
for i, result in enumerate(results[:5]):
|
||||||
print(f"{i+1}. '{result['word']}' at {result['position']}: {result['score']:.4f} (Δ: {result['improvement']:.4f}, tokens: {result['tokens']}, combined: {result['combined_score']:.4f})")
|
print(
|
||||||
|
f"{i+1}. '{result['word']}' at {result['position']}: {result['score']:.4f} (Δ: {result['improvement']:.4f}, tokens: {result['tokens']}, combined: {result['combined_score']:.4f})"
|
||||||
|
)
|
||||||
|
|
||||||
if best_result_idx >= 0:
|
if best_result_idx >= 0:
|
||||||
best_result = all_candidate_prefixes[best_result_idx]
|
best_result = all_candidate_prefixes[best_result_idx]
|
||||||
best_word = best_result["word"]
|
best_word = best_result["word"]
|
||||||
best_position = best_result["position"]
|
best_position = best_result["position"]
|
||||||
best_improvement = benign_scores[best_result_idx] - baseline_score
|
best_improvement = benign_scores[best_result_idx] - baseline_score
|
||||||
best_prefix = best_result["prefix"]
|
best_prefix = best_result["prefix"]
|
||||||
|
|
||||||
print(f"\nBest word to add: '{best_word}' at {best_position}")
|
print(f"\nBest word to add: '{best_word}' at {best_position}")
|
||||||
print(f"Improvement: {best_improvement:.4f} (from {baseline_score:.4f} to {benign_scores[best_result_idx]:.4f})")
|
print(
|
||||||
|
f"Improvement: {best_improvement:.4f} (from {baseline_score:.4f} to {benign_scores[best_result_idx]:.4f})"
|
||||||
|
)
|
||||||
print(f"New prefix: '{best_prefix}'")
|
print(f"New prefix: '{best_prefix}'")
|
||||||
return best_prefix, best_improvement
|
return best_prefix, best_improvement
|
||||||
else:
|
else:
|
||||||
print("No improvement found from any candidate word")
|
print("No improvement found from any candidate word")
|
||||||
return None, 0
|
return None, 0
|
||||||
|
|
||||||
|
|
||||||
def token_gradients_combined(
|
def token_gradients_combined(
|
||||||
model: AutoModelForSequenceClassification,
|
model: AutoModelForSequenceClassification,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
input_slice: slice,
|
input_slice: slice,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
benign_class: int = 1,
|
benign_class: int = 1,
|
||||||
malicious_class: int = 0,
|
malicious_class: int = 0,
|
||||||
alpha: float = 0.5,
|
alpha: float = 0.5,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
|
|
@ -257,12 +273,12 @@ def token_gradients_combined(
|
||||||
input_ids[input_slice].shape[0],
|
input_ids[input_slice].shape[0],
|
||||||
embed_weights.shape[0],
|
embed_weights.shape[0],
|
||||||
device=device,
|
device=device,
|
||||||
dtype=embed_weights.dtype
|
dtype=embed_weights.dtype,
|
||||||
)
|
)
|
||||||
one_hot.scatter_(
|
one_hot.scatter_(
|
||||||
1,
|
1,
|
||||||
input_ids[input_slice].unsqueeze(1),
|
input_ids[input_slice].unsqueeze(1),
|
||||||
torch.ones(one_hot.shape[0], 1, device=device, dtype=embed_weights.dtype)
|
torch.ones(one_hot.shape[0], 1, device=device, dtype=embed_weights.dtype),
|
||||||
)
|
)
|
||||||
one_hot.requires_grad_()
|
one_hot.requires_grad_()
|
||||||
input_embeds: torch.Tensor = (one_hot @ embed_weights).unsqueeze(0)
|
input_embeds: torch.Tensor = (one_hot @ embed_weights).unsqueeze(0)
|
||||||
|
|
@ -271,15 +287,18 @@ def token_gradients_combined(
|
||||||
embeds: torch.Tensor = model.deberta.embeddings.word_embeddings(input_ids)
|
embeds: torch.Tensor = model.deberta.embeddings.word_embeddings(input_ids)
|
||||||
full_embeds: torch.Tensor = torch.cat(
|
full_embeds: torch.Tensor = torch.cat(
|
||||||
[
|
[
|
||||||
embeds[:input_slice.start,:],
|
embeds[: input_slice.start, :],
|
||||||
input_embeds.squeeze(),
|
input_embeds.squeeze(),
|
||||||
embeds[input_slice.stop:,:]
|
embeds[input_slice.stop :, :],
|
||||||
],
|
],
|
||||||
dim=0)
|
dim=0,
|
||||||
|
)
|
||||||
logits: torch.Tensor = model(inputs_embeds=full_embeds.unsqueeze(0)).logits
|
logits: torch.Tensor = model(inputs_embeds=full_embeds.unsqueeze(0)).logits
|
||||||
|
|
||||||
# Combined loss: minimize malicious class (standard loss) and maximize benign class
|
# Combined loss: minimize malicious class (standard loss) and maximize benign class
|
||||||
standard_loss: torch.Tensor = nn.CrossEntropyLoss()(logits, torch.zeros(logits.shape[0], device=device).long())
|
standard_loss: torch.Tensor = nn.CrossEntropyLoss()(
|
||||||
|
logits, torch.zeros(logits.shape[0], device=device).long()
|
||||||
|
)
|
||||||
|
|
||||||
# Maximize benign class probability
|
# Maximize benign class probability
|
||||||
log_probs: torch.Tensor = torch.log_softmax(logits, dim=1)
|
log_probs: torch.Tensor = torch.log_softmax(logits, dim=1)
|
||||||
|
|
@ -293,75 +312,79 @@ def token_gradients_combined(
|
||||||
|
|
||||||
|
|
||||||
def analyze_token_contributions(
|
def analyze_token_contributions(
|
||||||
model: AutoModelForSequenceClassification,
|
model: AutoModelForSequenceClassification,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
injection_text: str,
|
injection_text: str,
|
||||||
adv_prefix: str,
|
adv_prefix: str,
|
||||||
text: str,
|
text: str,
|
||||||
benign_class_idx: int,
|
benign_class_idx: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
min_acceptable_benign: float = 0.6,
|
min_acceptable_benign: float = 0.6,
|
||||||
order_template: str = "{injection}{prefix}{text}" # Template for ordering components
|
order_template: str = "{injection}{prefix}{text}", # Template for ordering components
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Simple, non-batched approach to remove as many tokens as possible while keeping
|
Simple, non-batched approach to remove as many tokens as possible while keeping
|
||||||
the benign score above the minimum acceptable threshold.
|
the benign score above the minimum acceptable threshold.
|
||||||
"""
|
"""
|
||||||
print("\n----- ANALYZING TOKEN CONTRIBUTIONS (NO BATCHING) -----")
|
print("\n----- ANALYZING TOKEN CONTRIBUTIONS (NO BATCHING) -----")
|
||||||
|
|
||||||
# Get baseline benign score
|
# Get baseline benign score
|
||||||
full_text = order_template.format(injection=injection_text, prefix=adv_prefix, text=text)
|
full_text = order_template.format(injection=injection_text, prefix=adv_prefix, text=text)
|
||||||
inputs = tokenizer(full_text, return_tensors="pt")
|
inputs = tokenizer(full_text, return_tensors="pt")
|
||||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
inputs = {k: v.to(device) for k, v in inputs.items()}
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
logits = model(**inputs).logits
|
logits = model(**inputs).logits
|
||||||
probs = torch.softmax(logits, dim=-1)
|
probs = torch.softmax(logits, dim=-1)
|
||||||
baseline_score = probs[0][benign_class_idx].item()
|
baseline_score = probs[0][benign_class_idx].item()
|
||||||
|
|
||||||
print(f"Original prefix: '{adv_prefix}'")
|
print(f"Original prefix: '{adv_prefix}'")
|
||||||
print(f"Original benign score: {baseline_score:.4f}")
|
print(f"Original benign score: {baseline_score:.4f}")
|
||||||
|
|
||||||
# Use exactly the min_acceptable_benign as threshold
|
# Use exactly the min_acceptable_benign as threshold
|
||||||
threshold = min_acceptable_benign
|
threshold = min_acceptable_benign
|
||||||
print(f"Using threshold: {threshold:.4f}")
|
print(f"Using threshold: {threshold:.4f}")
|
||||||
|
|
||||||
if baseline_score < threshold:
|
if baseline_score < threshold:
|
||||||
print(f"Baseline score {baseline_score:.4f} already below threshold {threshold:.4f}. Stopping.")
|
print(
|
||||||
|
f"Baseline score {baseline_score:.4f} already below threshold {threshold:.4f}. Stopping."
|
||||||
|
)
|
||||||
return adv_prefix
|
return adv_prefix
|
||||||
|
|
||||||
current_prefix = adv_prefix
|
current_prefix = adv_prefix
|
||||||
remaining_tokens = tokenizer.tokenize(current_prefix)
|
remaining_tokens = tokenizer.tokenize(current_prefix)
|
||||||
print(f"Starting with {len(remaining_tokens)} tokens")
|
print(f"Starting with {len(remaining_tokens)} tokens")
|
||||||
|
|
||||||
removed_tokens = []
|
removed_tokens = []
|
||||||
|
|
||||||
while len(remaining_tokens) > 1:
|
while len(remaining_tokens) > 1:
|
||||||
# Try removing each token
|
# Try removing each token
|
||||||
best_candidate = None
|
best_candidate = None
|
||||||
best_score = -float('inf')
|
best_score = -float("inf")
|
||||||
best_idx = -1
|
best_idx = -1
|
||||||
|
|
||||||
for i in range(len(remaining_tokens)):
|
for i in range(len(remaining_tokens)):
|
||||||
# Create a new candidate with this token removed
|
# Create a new candidate with this token removed
|
||||||
tokens_without_i = remaining_tokens.copy()
|
tokens_without_i = remaining_tokens.copy()
|
||||||
token_to_remove = tokens_without_i.pop(i)
|
token_to_remove = tokens_without_i.pop(i)
|
||||||
candidate_prefix = tokenizer.convert_tokens_to_string(tokens_without_i)
|
candidate_prefix = tokenizer.convert_tokens_to_string(tokens_without_i)
|
||||||
|
|
||||||
# Evaluate this candidate
|
# Evaluate this candidate
|
||||||
full_text = order_template.format(injection=injection_text, prefix=candidate_prefix, text=text)
|
full_text = order_template.format(
|
||||||
|
injection=injection_text, prefix=candidate_prefix, text=text
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
inputs = tokenizer(full_text, return_tensors="pt")
|
inputs = tokenizer(full_text, return_tensors="pt")
|
||||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
inputs = {k: v.to(device) for k, v in inputs.items()}
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
logits = model(**inputs).logits
|
logits = model(**inputs).logits
|
||||||
probs = torch.softmax(logits, dim=-1)
|
probs = torch.softmax(logits, dim=-1)
|
||||||
score = probs[0][benign_class_idx].item()
|
score = probs[0][benign_class_idx].item()
|
||||||
|
|
||||||
print(f" Without token {i} ('{token_to_remove}'): score = {score:.4f}")
|
print(f" Without token {i} ('{token_to_remove}'): score = {score:.4f}")
|
||||||
|
|
||||||
# If this is still above threshold and better than our current best
|
# If this is still above threshold and better than our current best
|
||||||
if score >= threshold and score > best_score:
|
if score >= threshold and score > best_score:
|
||||||
best_candidate = candidate_prefix
|
best_candidate = candidate_prefix
|
||||||
|
|
@ -370,17 +393,19 @@ def analyze_token_contributions(
|
||||||
best_token = token_to_remove
|
best_token = token_to_remove
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f" Error evaluating without token {i}: {e}")
|
print(f" Error evaluating without token {i}: {e}")
|
||||||
|
|
||||||
# If we found a valid candidate, update our prefix
|
# If we found a valid candidate, update our prefix
|
||||||
if best_candidate:
|
if best_candidate:
|
||||||
current_prefix = best_candidate
|
current_prefix = best_candidate
|
||||||
removed_token = remaining_tokens.pop(best_idx)
|
removed_token = remaining_tokens.pop(best_idx)
|
||||||
removed_tokens.append(removed_token)
|
removed_tokens.append(removed_token)
|
||||||
print(f"✓ Removed token {best_idx} ('{best_token}'): new score = {best_score:.4f}, tokens left: {len(remaining_tokens)}")
|
print(
|
||||||
|
f"✓ Removed token {best_idx} ('{best_token}'): new score = {best_score:.4f}, tokens left: {len(remaining_tokens)}"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
print(f"Cannot remove any more tokens while staying above threshold {threshold:.4f}")
|
print(f"Cannot remove any more tokens while staying above threshold {threshold:.4f}")
|
||||||
break
|
break
|
||||||
|
|
||||||
# Final results
|
# Final results
|
||||||
print("\n===== TOKEN REMOVAL COMPLETE =====")
|
print("\n===== TOKEN REMOVAL COMPLETE =====")
|
||||||
print(f"Original prefix: '{adv_prefix}'")
|
print(f"Original prefix: '{adv_prefix}'")
|
||||||
|
|
@ -388,28 +413,29 @@ def analyze_token_contributions(
|
||||||
print(f"Removed {len(removed_tokens)} tokens: {removed_tokens}")
|
print(f"Removed {len(removed_tokens)} tokens: {removed_tokens}")
|
||||||
print(f"Original token count: {len(tokenizer.tokenize(adv_prefix))}")
|
print(f"Original token count: {len(tokenizer.tokenize(adv_prefix))}")
|
||||||
print(f"Final token count: {len(remaining_tokens)}")
|
print(f"Final token count: {len(remaining_tokens)}")
|
||||||
|
|
||||||
# Final verification
|
# Final verification
|
||||||
full_text = order_template.format(injection=injection_text, prefix=current_prefix, text=text)
|
full_text = order_template.format(injection=injection_text, prefix=current_prefix, text=text)
|
||||||
inputs = tokenizer(full_text, return_tensors="pt")
|
inputs = tokenizer(full_text, return_tensors="pt")
|
||||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
inputs = {k: v.to(device) for k, v in inputs.items()}
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
logits = model(**inputs).logits
|
logits = model(**inputs).logits
|
||||||
probs = torch.softmax(logits, dim=-1)
|
probs = torch.softmax(logits, dim=-1)
|
||||||
final_score = probs[0][benign_class_idx].item()
|
final_score = probs[0][benign_class_idx].item()
|
||||||
|
|
||||||
print(f"Final benign score: {final_score:.4f}")
|
print(f"Final benign score: {final_score:.4f}")
|
||||||
|
|
||||||
return current_prefix
|
return current_prefix
|
||||||
|
|
||||||
|
|
||||||
def minimize_tokens(
|
def minimize_tokens(
|
||||||
model: AutoModelForSequenceClassification,
|
model: AutoModelForSequenceClassification,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
injection_text: str,
|
injection_text: str,
|
||||||
adv_prefix: str,
|
adv_prefix: str,
|
||||||
text: str,
|
text: str,
|
||||||
benign_class_idx: int,
|
benign_class_idx: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
min_acceptable_benign: float = 0.6,
|
min_acceptable_benign: float = 0.6,
|
||||||
) -> str:
|
) -> str:
|
||||||
|
|
@ -417,85 +443,94 @@ def minimize_tokens(
|
||||||
Minimize tokens using only token contribution analysis (ablation study).
|
Minimize tokens using only token contribution analysis (ablation study).
|
||||||
This approach systematically removes tokens that contribute least to the benign classification.
|
This approach systematically removes tokens that contribute least to the benign classification.
|
||||||
Prioritizes removing shorter tokens when they have similar impacts on benign score.
|
Prioritizes removing shorter tokens when they have similar impacts on benign score.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
-----------
|
-----------
|
||||||
token_length_weight: Weight for prioritizing removal of short tokens (0-1, higher = prioritize short tokens more)
|
token_length_weight: Weight for prioritizing removal of short tokens (0-1, higher = prioritize short tokens more)
|
||||||
"""
|
"""
|
||||||
print("\n===== STARTING TOKEN MINIMIZATION =====")
|
print("\n===== STARTING TOKEN MINIMIZATION =====")
|
||||||
|
|
||||||
# Use only token ablation approach - systematically remove tokens that contribute least
|
# Use only token ablation approach - systematically remove tokens that contribute least
|
||||||
ablation_prefix: str = analyze_token_contributions(
|
ablation_prefix: str = analyze_token_contributions(
|
||||||
model, tokenizer, injection_text, adv_prefix, text,
|
model,
|
||||||
|
tokenizer,
|
||||||
|
injection_text,
|
||||||
|
adv_prefix,
|
||||||
|
text,
|
||||||
benign_class_idx,
|
benign_class_idx,
|
||||||
device=device, min_acceptable_benign=min_acceptable_benign,
|
device=device,
|
||||||
|
min_acceptable_benign=min_acceptable_benign,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Report final token count
|
# Report final token count
|
||||||
final_token_count: int = len(tokenizer.encode(ablation_prefix, add_special_tokens=False))
|
final_token_count: int = len(tokenizer.encode(ablation_prefix, add_special_tokens=False))
|
||||||
original_token_count: int = len(tokenizer.encode(adv_prefix, add_special_tokens=False))
|
original_token_count: int = len(tokenizer.encode(adv_prefix, add_special_tokens=False))
|
||||||
|
|
||||||
print(f"\n===== TOKEN MINIMIZATION COMPLETE =====")
|
print(f"\n===== TOKEN MINIMIZATION COMPLETE =====")
|
||||||
print(f"Original token count: {original_token_count}")
|
print(f"Original token count: {original_token_count}")
|
||||||
print(f"Final token count: {final_token_count}")
|
print(f"Final token count: {final_token_count}")
|
||||||
print(f"Reduction: {((original_token_count - final_token_count) / original_token_count * 100):.2f}%")
|
print(
|
||||||
|
f"Reduction: {((original_token_count - final_token_count) / original_token_count * 100):.2f}%"
|
||||||
|
)
|
||||||
print(f"Final prefix: '{ablation_prefix}'")
|
print(f"Final prefix: '{ablation_prefix}'")
|
||||||
|
|
||||||
return ablation_prefix
|
return ablation_prefix
|
||||||
|
|
||||||
|
|
||||||
def sample_control(
|
def sample_control(
|
||||||
control_toks: torch.Tensor,
|
control_toks: torch.Tensor,
|
||||||
grad: torch.Tensor,
|
grad: torch.Tensor,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
topk: int = 256,
|
topk: int = 256,
|
||||||
temp: float = 1,
|
temp: float = 1,
|
||||||
not_allowed_tokens: Optional[torch.Tensor] = None
|
not_allowed_tokens: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if not_allowed_tokens is not None:
|
if not_allowed_tokens is not None:
|
||||||
grad[:, not_allowed_tokens.to(grad.device)] = float('inf')
|
grad[:, not_allowed_tokens.to(grad.device)] = float("inf")
|
||||||
|
|
||||||
top_indices: torch.Tensor = (-grad).topk(topk, dim=1).indices
|
top_indices: torch.Tensor = (-grad).topk(topk, dim=1).indices
|
||||||
control_toks = control_toks.to(grad.device)
|
control_toks = control_toks.to(grad.device)
|
||||||
|
|
||||||
original_control_toks: torch.Tensor = control_toks.repeat(batch_size, 1)
|
original_control_toks: torch.Tensor = control_toks.repeat(batch_size, 1)
|
||||||
|
|
||||||
# Ensure batch_size doesn't exceed the size of control_toks
|
# Ensure batch_size doesn't exceed the size of control_toks
|
||||||
actual_batch_size: int = min(batch_size, len(control_toks))
|
actual_batch_size: int = min(batch_size, len(control_toks))
|
||||||
|
|
||||||
new_token_pos: torch.Tensor = torch.arange(
|
new_token_pos: torch.Tensor = torch.arange(
|
||||||
0,
|
0,
|
||||||
len(control_toks),
|
len(control_toks),
|
||||||
max(1, len(control_toks) / actual_batch_size), # Ensure step is at least 1
|
max(1, len(control_toks) / actual_batch_size), # Ensure step is at least 1
|
||||||
device=grad.device
|
device=grad.device,
|
||||||
).type(torch.int64)
|
).type(torch.int64)
|
||||||
|
|
||||||
# Extra safety: ensure new_token_pos is within bounds of top_indices' first dimension
|
# Extra safety: ensure new_token_pos is within bounds of top_indices' first dimension
|
||||||
new_token_pos = torch.clamp(new_token_pos, 0, grad.shape[0] - 1)
|
new_token_pos = torch.clamp(new_token_pos, 0, grad.shape[0] - 1)
|
||||||
|
|
||||||
new_token_val: torch.Tensor = torch.gather(
|
new_token_val: torch.Tensor = torch.gather(
|
||||||
top_indices[new_token_pos], 1,
|
top_indices[new_token_pos],
|
||||||
torch.randint(0, topk, (len(new_token_pos), 1), device=grad.device)
|
1,
|
||||||
|
torch.randint(0, topk, (len(new_token_pos), 1), device=grad.device),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Ensure we don't exceed the original batch size dimension
|
# Ensure we don't exceed the original batch size dimension
|
||||||
new_control_toks: torch.Tensor = original_control_toks[:len(new_token_pos)].scatter_(
|
new_control_toks: torch.Tensor = original_control_toks[: len(new_token_pos)].scatter_(
|
||||||
1, new_token_pos.unsqueeze(-1), new_token_val
|
1, new_token_pos.unsqueeze(-1), new_token_val
|
||||||
)
|
)
|
||||||
|
|
||||||
return new_control_toks
|
return new_control_toks
|
||||||
|
|
||||||
|
|
||||||
def get_random_words(n: int = 10, min_uses: int = 0, token_priority: float = 0.3) -> List[str]:
|
def get_random_words(n: int = 10, min_uses: int = 0, token_priority: float = 0.3) -> List[str]:
|
||||||
"""
|
"""
|
||||||
Get a list of words to use, prioritizing words that have performed well in the past.
|
Get a list of words to use, prioritizing words that have performed well in the past.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
-----------
|
-----------
|
||||||
n: Number of words to return
|
n: Number of words to return
|
||||||
min_uses: Minimum number of uses a word must have to be considered from the database
|
min_uses: Minimum number of uses a word must have to be considered from the database
|
||||||
token_priority: How much to prioritize words with fewer tokens (0-1)
|
token_priority: How much to prioritize words with fewer tokens (0-1)
|
||||||
0 = purely improvement based, 1 = purely token count based
|
0 = purely improvement based, 1 = purely token count based
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
--------
|
--------
|
||||||
List of words
|
List of words
|
||||||
|
|
@ -510,20 +545,18 @@ def get_random_words(n: int = 10, min_uses: int = 0, token_priority: float = 0.3
|
||||||
else:
|
else:
|
||||||
# Use combined sorting with the specified token weight
|
# Use combined sorting with the specified token weight
|
||||||
top_words = words_db.get_top_words(
|
top_words = words_db.get_top_words(
|
||||||
limit=n,
|
limit=n, min_uses=min_uses, sort_by="combined", token_weight=token_priority
|
||||||
min_uses=min_uses,
|
|
||||||
sort_by="combined",
|
|
||||||
token_weight=token_priority
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# If we got enough words from the database, use them
|
# If we got enough words from the database, use them
|
||||||
if len(top_words) >= n:
|
if len(top_words) >= n:
|
||||||
return top_words[:n]
|
return top_words[:n]
|
||||||
|
|
||||||
# Otherwise, use what we got plus some random words
|
# Otherwise, use what we got plus some random words
|
||||||
remaining = n - len(top_words)
|
remaining = n - len(top_words)
|
||||||
return top_words + random.choices(words, k=remaining)
|
return top_words + random.choices(words, k=remaining)
|
||||||
|
|
||||||
|
|
||||||
def count_tokens(text: str, model: str = "gpt-3.5") -> int:
|
def count_tokens(text: str, model: str = "gpt-3.5") -> int:
|
||||||
"""Count the number of tokens in a text string using tiktoken."""
|
"""Count the number of tokens in a text string using tiktoken."""
|
||||||
try:
|
try:
|
||||||
|
|
@ -533,14 +566,15 @@ def count_tokens(text: str, model: str = "gpt-3.5") -> int:
|
||||||
# Fallback to a simple approximation if tiktoken fails
|
# Fallback to a simple approximation if tiktoken fails
|
||||||
return len(text.split())
|
return len(text.split())
|
||||||
|
|
||||||
|
|
||||||
def get_combined_score(
|
def get_combined_score(
|
||||||
model: AutoModelForSequenceClassification,
|
model: AutoModelForSequenceClassification,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
text: str,
|
text: str,
|
||||||
candidates: List[str],
|
candidates: List[str],
|
||||||
benign_idx: int,
|
benign_idx: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
alpha: float = 0.5,
|
alpha: float = 0.5,
|
||||||
token_penalty_weight: float = 0.1,
|
token_penalty_weight: float = 0.1,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""
|
"""
|
||||||
|
|
@ -570,13 +604,13 @@ def get_combined_score(
|
||||||
int
|
int
|
||||||
Index of the candidate with the best combined score.
|
Index of the candidate with the best combined score.
|
||||||
"""
|
"""
|
||||||
best_combined_score: float = float('-inf')
|
best_combined_score: float = float("-inf")
|
||||||
best_idx: int = 0
|
best_idx: int = 0
|
||||||
|
|
||||||
# Find max token count for normalization
|
# Find max token count for normalization
|
||||||
token_counts: List[int] = [count_tokens(candidate) for candidate in candidates]
|
token_counts: List[int] = [count_tokens(candidate) for candidate in candidates]
|
||||||
max_token_count: int = max(token_counts) if token_counts else 1
|
max_token_count: int = max(token_counts) if token_counts else 1
|
||||||
|
|
||||||
for idx, candidate in enumerate(candidates):
|
for idx, candidate in enumerate(candidates):
|
||||||
inputs: Dict[str, torch.Tensor] = tokenizer(candidate + text, return_tensors="pt")
|
inputs: Dict[str, torch.Tensor] = tokenizer(candidate + text, return_tensors="pt")
|
||||||
# Move inputs to MPS device
|
# Move inputs to MPS device
|
||||||
|
|
@ -586,21 +620,27 @@ def get_combined_score(
|
||||||
probs: torch.Tensor = torch.softmax(logits, dim=-1)
|
probs: torch.Tensor = torch.softmax(logits, dim=-1)
|
||||||
|
|
||||||
# Calculate loss component (lower is better)
|
# Calculate loss component (lower is better)
|
||||||
loss: torch.Tensor = nn.CrossEntropyLoss()(logits, torch.zeros(logits.shape[0], device=device).long())
|
loss: torch.Tensor = nn.CrossEntropyLoss()(
|
||||||
normalized_loss: float = 1.0 / (1.0 + loss.item()) # Convert to 0-1 range where higher is better
|
logits, torch.zeros(logits.shape[0], device=device).long()
|
||||||
|
)
|
||||||
|
normalized_loss: float = 1.0 / (
|
||||||
|
1.0 + loss.item()
|
||||||
|
) # Convert to 0-1 range where higher is better
|
||||||
|
|
||||||
# Calculate benign score component (higher is better)
|
# Calculate benign score component (higher is better)
|
||||||
benign_score: float = probs[0][benign_idx].item()
|
benign_score: float = probs[0][benign_idx].item()
|
||||||
|
|
||||||
# Calculate token count penalty (normalized to 0-1, where higher is better = fewer tokens)
|
# Calculate token count penalty (normalized to 0-1, where higher is better = fewer tokens)
|
||||||
token_count: int = token_counts[idx]
|
token_count: int = token_counts[idx]
|
||||||
token_penalty: float = 1.0 - (token_count / max_token_count)
|
token_penalty: float = 1.0 - (token_count / max_token_count)
|
||||||
|
|
||||||
# Calculate combined score with token penalty
|
# Calculate combined score with token penalty
|
||||||
combined_score: float = (alpha * benign_score + (1 - alpha) * normalized_loss) * (1 - token_penalty_weight + token_penalty_weight * token_penalty)
|
combined_score: float = (alpha * benign_score + (1 - alpha) * normalized_loss) * (
|
||||||
|
1 - token_penalty_weight + token_penalty_weight * token_penalty
|
||||||
|
)
|
||||||
|
|
||||||
if combined_score > best_combined_score:
|
if combined_score > best_combined_score:
|
||||||
best_combined_score = combined_score
|
best_combined_score = combined_score
|
||||||
best_idx = idx
|
best_idx = idx
|
||||||
|
|
||||||
return best_idx
|
return best_idx
|
||||||
|
|
|
||||||
609
words.py
609
words.py
|
|
@ -1,469 +1,3 @@
|
||||||
words = ["Please",
|
|
||||||
"Thank you",
|
|
||||||
"Excuse me",
|
|
||||||
"Pardon me",
|
|
||||||
"Sorry",
|
|
||||||
"Welcome",
|
|
||||||
"Certainly",
|
|
||||||
"Absolutely",
|
|
||||||
"Indeed",
|
|
||||||
"Agreed",
|
|
||||||
"Agreed",
|
|
||||||
"Agreed",
|
|
||||||
"Okay",
|
|
||||||
"Alright",
|
|
||||||
"Very well",
|
|
||||||
"Of course",
|
|
||||||
"By all means",
|
|
||||||
"No problem",
|
|
||||||
"You're welcome",
|
|
||||||
"My pleasure",
|
|
||||||
"Don't mention it",
|
|
||||||
"Anytime",
|
|
||||||
"With pleasure",
|
|
||||||
"Gladly",
|
|
||||||
"Happily",
|
|
||||||
"Kindly",
|
|
||||||
"Gently",
|
|
||||||
"Softly",
|
|
||||||
"Patiently",
|
|
||||||
"Carefully",
|
|
||||||
"Thoughtfully",
|
|
||||||
"Considerately",
|
|
||||||
"Respectfully",
|
|
||||||
"Humbly",
|
|
||||||
"Gratefully",
|
|
||||||
"Thankfully",
|
|
||||||
"Appreciatively",
|
|
||||||
"Sincerely",
|
|
||||||
"Truly",
|
|
||||||
"Genuinely",
|
|
||||||
"Honestly",
|
|
||||||
"Frankly",
|
|
||||||
"Openly",
|
|
||||||
"Candidly",
|
|
||||||
"Politely",
|
|
||||||
"Courteously",
|
|
||||||
"Graciously",
|
|
||||||
"Charmingly",
|
|
||||||
"Amiably",
|
|
||||||
"Genially",
|
|
||||||
"Cordially",
|
|
||||||
"Warmly",
|
|
||||||
"Friendly",
|
|
||||||
"Welcoming",
|
|
||||||
"Inviting",
|
|
||||||
"Pleasant",
|
|
||||||
"Agreeable",
|
|
||||||
"Kind",
|
|
||||||
"Nice",
|
|
||||||
"Sweet",
|
|
||||||
"Lovely",
|
|
||||||
"Delightful",
|
|
||||||
"Wonderful",
|
|
||||||
"Excellent",
|
|
||||||
"Great",
|
|
||||||
"Fantastic",
|
|
||||||
"Amazing",
|
|
||||||
"Superb",
|
|
||||||
"Brilliant",
|
|
||||||
"Splendid",
|
|
||||||
"Marvelous",
|
|
||||||
"Terrific",
|
|
||||||
"Awesome",
|
|
||||||
"Fabulous",
|
|
||||||
"Spectacular",
|
|
||||||
"Stupendous",
|
|
||||||
"Phenomenal",
|
|
||||||
"Remarkable",
|
|
||||||
"Impressive",
|
|
||||||
"Admirable",
|
|
||||||
"Commendable",
|
|
||||||
"Praiseworthy",
|
|
||||||
"Respectable",
|
|
||||||
"Honorable",
|
|
||||||
"Dignified",
|
|
||||||
"Noble",
|
|
||||||
"Benevolent",
|
|
||||||
"Generous",
|
|
||||||
"Charitable",
|
|
||||||
"Giving",
|
|
||||||
"Helpful",
|
|
||||||
"Cooperative",
|
|
||||||
"Accommodating",
|
|
||||||
"Obliging",
|
|
||||||
"Supportive",
|
|
||||||
"Understanding",
|
|
||||||
"Empathetic",
|
|
||||||
"Compassionate",
|
|
||||||
"Caring",
|
|
||||||
"Loving"]
|
|
||||||
|
|
||||||
words2 = [
|
|
||||||
"Please",
|
|
||||||
"Thanks",
|
|
||||||
"Sorry",
|
|
||||||
"Excuse",
|
|
||||||
"Pardon",
|
|
||||||
"Welcome",
|
|
||||||
"Kindly",
|
|
||||||
"May",
|
|
||||||
"Could",
|
|
||||||
"Would",
|
|
||||||
"Shall",
|
|
||||||
"Might",
|
|
||||||
"Do",
|
|
||||||
"Certainly",
|
|
||||||
"Indeed",
|
|
||||||
"Absolutely",
|
|
||||||
"Definitely",
|
|
||||||
"Naturally",
|
|
||||||
"Precisely",
|
|
||||||
"Assuredly",
|
|
||||||
"Undoubtedly",
|
|
||||||
"Gladly",
|
|
||||||
"Sure",
|
|
||||||
"Alright",
|
|
||||||
"Okay",
|
|
||||||
"OK",
|
|
||||||
"Fine",
|
|
||||||
"Fair",
|
|
||||||
"Aye",
|
|
||||||
"Yea",
|
|
||||||
"Obliged",
|
|
||||||
"Sir",
|
|
||||||
"Madam",
|
|
||||||
"Ma'am",
|
|
||||||
"Well",
|
|
||||||
"Ah",
|
|
||||||
"Oh",
|
|
||||||
"Just",
|
|
||||||
"Good",
|
|
||||||
"Permit",
|
|
||||||
"Allow",
|
|
||||||
"Grant",
|
|
||||||
"Proffer",
|
|
||||||
"Beg",
|
|
||||||
"Trust",
|
|
||||||
"Hope",
|
|
||||||
"Greetings",
|
|
||||||
"Salutations",
|
|
||||||
"Regards",
|
|
||||||
"Congrats",
|
|
||||||
"Congratulations",
|
|
||||||
"Bravo",
|
|
||||||
"Kudos",
|
|
||||||
"Farewell",
|
|
||||||
"Adieu",
|
|
||||||
"Ciao",
|
|
||||||
"Gracious",
|
|
||||||
"Mercy",
|
|
||||||
"Bless",
|
|
||||||
"Pray",
|
|
||||||
"Prithee",
|
|
||||||
"Hark",
|
|
||||||
"Henceforth",
|
|
||||||
"Henceforward",
|
|
||||||
"Hence",
|
|
||||||
"Forsooth",
|
|
||||||
"Respectfully",
|
|
||||||
"Sincerely",
|
|
||||||
"Truly",
|
|
||||||
"Frankly",
|
|
||||||
"Honestly",
|
|
||||||
"Genuinely",
|
|
||||||
"Openly",
|
|
||||||
"Candidly",
|
|
||||||
"Politely",
|
|
||||||
"Courteously",
|
|
||||||
"Graciously",
|
|
||||||
]
|
|
||||||
|
|
||||||
words3 = [
|
|
||||||
"description",
|
|
||||||
"manifest",
|
|
||||||
"reddit",
|
|
||||||
"recruit",
|
|
||||||
"flight",
|
|
||||||
"check",
|
|
||||||
"position",
|
|
||||||
"respectfully",
|
|
||||||
"bless",
|
|
||||||
"generator",
|
|
||||||
"reading",
|
|
||||||
"grave",
|
|
||||||
"medicine",
|
|
||||||
"paper",
|
|
||||||
"cleaning",
|
|
||||||
"related",
|
|
||||||
"foul",
|
|
||||||
"width",
|
|
||||||
"characteristics",
|
|
||||||
"rotate",
|
|
||||||
"logistic",
|
|
||||||
"named",
|
|
||||||
"correction",
|
|
||||||
"select",
|
|
||||||
"consider",
|
|
||||||
"other",
|
|
||||||
"missing",
|
|
||||||
"advertising",
|
|
||||||
"named",
|
|
||||||
"inbound",
|
|
||||||
"rate",
|
|
||||||
"suicide",
|
|
||||||
"shortDescription",
|
|
||||||
"catcher",
|
|
||||||
"concurrent",
|
|
||||||
"chemistry",
|
|
||||||
"fighting",
|
|
||||||
"complain",
|
|
||||||
"score",
|
|
||||||
"downloading",
|
|
||||||
"medstation",
|
|
||||||
"Bangkok",
|
|
||||||
"missing",
|
|
||||||
"weebly",
|
|
||||||
"garnitur",
|
|
||||||
"sporto",
|
|
||||||
"cyclosporto",
|
|
||||||
"LinkedIn",
|
|
||||||
"basket",
|
|
||||||
"nut",
|
|
||||||
"Lifettc",
|
|
||||||
"Collect",
|
|
||||||
"stonk",
|
|
||||||
"vinner",
|
|
||||||
"rønde",
|
|
||||||
"Collect",
|
|
||||||
"iris",
|
|
||||||
"Simon",
|
|
||||||
"cleaning",
|
|
||||||
"related",
|
|
||||||
"သာသနာ",
|
|
||||||
"theatre",
|
|
||||||
"gemaak",
|
|
||||||
"куча",
|
|
||||||
"народ",
|
|
||||||
"correction",
|
|
||||||
"Bang",
|
|
||||||
"category",
|
|
||||||
"catcher",
|
|
||||||
"参照",
|
|
||||||
"separate",
|
|
||||||
"almal",
|
|
||||||
"Bangkok",
|
|
||||||
"missing",
|
|
||||||
"stock",
|
|
||||||
"youtube",
|
|
||||||
"attention",
|
|
||||||
"fighting",
|
|
||||||
"respectfully",
|
|
||||||
"Place",
|
|
||||||
"Upload",
|
|
||||||
"next",
|
|
||||||
"words",
|
|
||||||
"Moi",
|
|
||||||
"NAMA",
|
|
||||||
"mandar",
|
|
||||||
"alquiler",
|
|
||||||
"chat",
|
|
||||||
"Sebab",
|
|
||||||
"Perfect",
|
|
||||||
"distinct",
|
|
||||||
"bots",
|
|
||||||
"Ing",
|
|
||||||
"falt",
|
|
||||||
"placements",
|
|
||||||
"sivo",
|
|
||||||
"else",
|
|
||||||
"はお",
|
|
||||||
"ICA",
|
|
||||||
"Цвет",
|
|
||||||
"Check",
|
|
||||||
"valid",
|
|
||||||
"earn",
|
|
||||||
"con",
|
|
||||||
"villa",
|
|
||||||
"outil",
|
|
||||||
"Sun",
|
|
||||||
"vertido",
|
|
||||||
"en",
|
|
||||||
"Dub",
|
|
||||||
"danza",
|
|
||||||
"Articolo",
|
|
||||||
"Vsions",
|
|
||||||
"Cruise",
|
|
||||||
"Saatchara",
|
|
||||||
"ала",
|
|
||||||
"source",
|
|
||||||
"ungalow",
|
|
||||||
"TITLE",
|
|
||||||
"gén",
|
|
||||||
"セكية",
|
|
||||||
"Fra",
|
|
||||||
"英会話",
|
|
||||||
"Verstaking",
|
|
||||||
"Just",
|
|
||||||
"Teacher",
|
|
||||||
"itelji",
|
|
||||||
"Hot",
|
|
||||||
"Palquis",
|
|
||||||
"enez",
|
|
||||||
"Man",
|
|
||||||
"Recommend",
|
|
||||||
"YouTube",
|
|
||||||
"attention",
|
|
||||||
"foulo",
|
|
||||||
"original",
|
|
||||||
"grave",
|
|
||||||
"May",
|
|
||||||
"compete",
|
|
||||||
"Metro",
|
|
||||||
"wacomercia",
|
|
||||||
"this",
|
|
||||||
"combat",
|
|
||||||
"verencolor",
|
|
||||||
"STAM",
|
|
||||||
"ilä",
|
|
||||||
"visit",
|
|
||||||
"toy",
|
|
||||||
"additional",
|
|
||||||
"在中国",
|
|
||||||
"cnhaben",
|
|
||||||
"same",
|
|
||||||
"including",
|
|
||||||
"term",
|
|
||||||
"注意到",
|
|
||||||
"position",
|
|
||||||
"Ingredients",
|
|
||||||
"classification",
|
|
||||||
"dimensions",
|
|
||||||
"REVIS",
|
|
||||||
"meteor",
|
|
||||||
"information",
|
|
||||||
"Term",
|
|
||||||
"giene",
|
|
||||||
"Teacher",
|
|
||||||
"Should",
|
|
||||||
"gala",
|
|
||||||
"부",
|
|
||||||
"mention",
|
|
||||||
"postal",
|
|
||||||
"foul",
|
|
||||||
"страница",
|
|
||||||
"respectfully",
|
|
||||||
"cutive",
|
|
||||||
"fighting",
|
|
||||||
"instrui",
|
|
||||||
"Songs",
|
|
||||||
"Christian",
|
|
||||||
"song",
|
|
||||||
"all",
|
|
||||||
"Мал",
|
|
||||||
"ozou",
|
|
||||||
"mus",
|
|
||||||
"bron",
|
|
||||||
"rhythm",
|
|
||||||
"əчитель",
|
|
||||||
"sis",
|
|
||||||
"tarra",
|
|
||||||
"Abdul",
|
|
||||||
"publish",
|
|
||||||
"consulta",
|
|
||||||
"amlustra",
|
|
||||||
"useful",
|
|
||||||
"classification",
|
|
||||||
"brief",
|
|
||||||
"Fall",
|
|
||||||
"amina",
|
|
||||||
"Carbon",
|
|
||||||
"bertso",
|
|
||||||
"Attend",
|
|
||||||
"licenses",
|
|
||||||
"sections",
|
|
||||||
"cidos",
|
|
||||||
"below",
|
|
||||||
"ículo",
|
|
||||||
"gehalt",
|
|
||||||
"alphabet",
|
|
||||||
"respective",
|
|
||||||
"izoh",
|
|
||||||
"plementar",
|
|
||||||
"mention",
|
|
||||||
"information",
|
|
||||||
"typeof",
|
|
||||||
"可见",
|
|
||||||
"giene",
|
|
||||||
"suits",
|
|
||||||
"sponso",
|
|
||||||
"gossip",
|
|
||||||
"Medical",
|
|
||||||
"er",
|
|
||||||
"ан",
|
|
||||||
"other",
|
|
||||||
"com",
|
|
||||||
"chemistry",
|
|
||||||
"recommend",
|
|
||||||
"description",
|
|
||||||
"vagina",
|
|
||||||
"注意到",
|
|
||||||
"album",
|
|
||||||
"hint",
|
|
||||||
"places",
|
|
||||||
"cart",
|
|
||||||
"celestial",
|
|
||||||
"prene",
|
|
||||||
"netter",
|
|
||||||
"yours",
|
|
||||||
"art",
|
|
||||||
"traduc",
|
|
||||||
"Optional",
|
|
||||||
"all",
|
|
||||||
"different",
|
|
||||||
"casi",
|
|
||||||
"Boas",
|
|
||||||
"estra",
|
|
||||||
"vis",
|
|
||||||
"client",
|
|
||||||
"ABLE",
|
|
||||||
"under",
|
|
||||||
"export",
|
|
||||||
"EVER",
|
|
||||||
"pro",
|
|
||||||
"exe",
|
|
||||||
"iliye",
|
|
||||||
"chemical",
|
|
||||||
"Дести",
|
|
||||||
"graphics",
|
|
||||||
"Bean",
|
|
||||||
"marques",
|
|
||||||
"kanta",
|
|
||||||
"talaga",
|
|
||||||
"Found",
|
|
||||||
"Sa",
|
|
||||||
"Html",
|
|
||||||
"transformation",
|
|
||||||
"extra",
|
|
||||||
"Wort",
|
|
||||||
"Male",
|
|
||||||
"Soda",
|
|
||||||
"drip",
|
|
||||||
"recommend",
|
|
||||||
"Friends",
|
|
||||||
"avoid",
|
|
||||||
"formant",
|
|
||||||
"type",
|
|
||||||
"font",
|
|
||||||
"UA",
|
|
||||||
"keyword",
|
|
||||||
"Term",
|
|
||||||
"Medical"
|
|
||||||
]
|
|
||||||
|
|
||||||
words3 = list(set(words3))
|
|
||||||
|
|
||||||
|
|
||||||
words4 = [
|
words4 = [
|
||||||
"ocular",
|
"ocular",
|
||||||
"spell",
|
"spell",
|
||||||
|
|
@ -491,19 +25,33 @@ words4 = [
|
||||||
"recommend",
|
"recommend",
|
||||||
"better",
|
"better",
|
||||||
"previous",
|
"previous",
|
||||||
"Camp",
|
"Camp",
|
||||||
"Megan",
|
"Megan",
|
||||||
"Voltage",
|
"Voltage",
|
||||||
"Previous",
|
"Previous",
|
||||||
"reviews",
|
"reviews",
|
||||||
"instru",
|
"instru",
|
||||||
"Exercise",
|
"Exercise",
|
||||||
"instru",
|
"instru",
|
||||||
"Russian",
|
"Russian",
|
||||||
"vaccination",
|
"vaccination",
|
||||||
"can","dread",
|
"can",
|
||||||
'all', 'Simon', 'material', 'spray', 'Subjects', 'recess', 'position', 'contrast', 'want', 'twenty', 'dependent', 'recommend', 'read',
|
"dread",
|
||||||
'sección', 'Hospital',
|
"all",
|
||||||
|
"Simon",
|
||||||
|
"material",
|
||||||
|
"spray",
|
||||||
|
"Subjects",
|
||||||
|
"recess",
|
||||||
|
"position",
|
||||||
|
"contrast",
|
||||||
|
"want",
|
||||||
|
"twenty",
|
||||||
|
"dependent",
|
||||||
|
"recommend",
|
||||||
|
"read",
|
||||||
|
"sección",
|
||||||
|
"Hospital",
|
||||||
"citation",
|
"citation",
|
||||||
"edge",
|
"edge",
|
||||||
"solid",
|
"solid",
|
||||||
|
|
@ -534,22 +82,21 @@ words4 = [
|
||||||
"今年",
|
"今年",
|
||||||
"许可证号",
|
"许可证号",
|
||||||
"nutrition",
|
"nutrition",
|
||||||
"previous",
|
"previous",
|
||||||
"additional",
|
"additional",
|
||||||
"better",
|
"better",
|
||||||
"Word",
|
"Word",
|
||||||
"leg",
|
"leg",
|
||||||
"similar",
|
"similar",
|
||||||
"anchors",
|
"anchors",
|
||||||
"civil",
|
"civil",
|
||||||
"civil",
|
"civil",
|
||||||
"are",
|
"are",
|
||||||
"bilar",
|
"bilar",
|
||||||
"calanus",
|
"calanus",
|
||||||
"ITIS",
|
"ITIS",
|
||||||
"量的",
|
"量的",
|
||||||
"1657"
|
"1657" "depression",
|
||||||
"depression",
|
|
||||||
"Upload",
|
"Upload",
|
||||||
"overseas",
|
"overseas",
|
||||||
"irrigate",
|
"irrigate",
|
||||||
|
|
@ -622,9 +169,22 @@ words4 = [
|
||||||
"question",
|
"question",
|
||||||
"faillets",
|
"faillets",
|
||||||
"耙もの",
|
"耙もの",
|
||||||
"prestencil", "vine", "birds", "help", "Container", "mention",
|
"prestencil",
|
||||||
"Primary", "participation", "Maintenance", "Categories", "malaysia",
|
"vine",
|
||||||
"vascular", "editorial", "OECD", "question", "consider",
|
"birds",
|
||||||
|
"help",
|
||||||
|
"Container",
|
||||||
|
"mention",
|
||||||
|
"Primary",
|
||||||
|
"participation",
|
||||||
|
"Maintenance",
|
||||||
|
"Categories",
|
||||||
|
"malaysia",
|
||||||
|
"vascular",
|
||||||
|
"editorial",
|
||||||
|
"OECD",
|
||||||
|
"question",
|
||||||
|
"consider",
|
||||||
"attachment",
|
"attachment",
|
||||||
"information",
|
"information",
|
||||||
"recommend",
|
"recommend",
|
||||||
|
|
@ -753,8 +313,7 @@ words4 = [
|
||||||
"question",
|
"question",
|
||||||
"consider",
|
"consider",
|
||||||
"oval",
|
"oval",
|
||||||
"preferred"
|
"preferred" "resources",
|
||||||
"resources",
|
|
||||||
"phrases",
|
"phrases",
|
||||||
"low",
|
"low",
|
||||||
"Mark",
|
"Mark",
|
||||||
|
|
@ -794,37 +353,47 @@ words4 = [
|
||||||
"form",
|
"form",
|
||||||
"similar",
|
"similar",
|
||||||
"candid",
|
"candid",
|
||||||
'Night', 'similar', 'atelier', 'keyword', 'repository', 'maintain', 'physique', 'excessopathy', 'article', 'information', 'recommend', 'consider'
|
"Night",
|
||||||
"irish",
|
"similar",
|
||||||
|
"atelier",
|
||||||
|
"keyword",
|
||||||
|
"repository",
|
||||||
|
"maintain",
|
||||||
|
"physique",
|
||||||
|
"excessopathy",
|
||||||
|
"article",
|
||||||
|
"information",
|
||||||
|
"recommend",
|
||||||
|
"consider" "irish",
|
||||||
"accessories",
|
"accessories",
|
||||||
"caption",
|
"caption",
|
||||||
"pression",
|
"pression",
|
||||||
"secteur", # French for sector
|
"secteur", # French for sector
|
||||||
"tag",
|
"tag",
|
||||||
"category",
|
"category",
|
||||||
"sebelum", # Indonesian for before
|
"sebelum", # Indonesian for before
|
||||||
"zoom",
|
"zoom",
|
||||||
"reibung", # German for friction
|
"reibung", # German for friction
|
||||||
"tension",
|
"tension",
|
||||||
"nutrient",
|
"nutrient",
|
||||||
"layer",
|
"layer",
|
||||||
"below",
|
"below",
|
||||||
"recommend",
|
"recommend",
|
||||||
'previous',
|
"previous",
|
||||||
'Brush',
|
"Brush",
|
||||||
'write',
|
"write",
|
||||||
'some',
|
"some",
|
||||||
'needle',
|
"needle",
|
||||||
'same',
|
"same",
|
||||||
'antioxidant',
|
"antioxidant",
|
||||||
'are',
|
"are",
|
||||||
'separate',
|
"separate",
|
||||||
'注意',
|
"注意",
|
||||||
'кула',
|
"кула",
|
||||||
'лист',
|
"лист",
|
||||||
'液压',
|
"液压",
|
||||||
'gène',
|
"gène",
|
||||||
'bel'
|
"bel",
|
||||||
]
|
]
|
||||||
|
|
||||||
words4 = list(set(words4[:230]))
|
words4 = list(set(words4[:230]))
|
||||||
|
|
|
||||||
157
wordsdb.py
157
wordsdb.py
|
|
@ -2,25 +2,28 @@ import sqlite3
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import List, Optional, Dict, Any
|
from typing import List, Optional, Dict, Any
|
||||||
|
|
||||||
|
|
||||||
class WordsDatabase:
|
class WordsDatabase:
|
||||||
"""
|
"""
|
||||||
Database to track the performance of words when added to a prefix.
|
Database to track the performance of words when added to a prefix.
|
||||||
Stores word statistics and allows querying for top-performing words.
|
Stores word statistics and allows querying for top-performing words.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, db_path: str = "word_performance.db"):
|
def __init__(self, db_path: str = "word_performance.db"):
|
||||||
"""Initialize the database, creating tables if they don't exist."""
|
"""Initialize the database, creating tables if they don't exist."""
|
||||||
self.db_path = db_path
|
self.db_path = db_path
|
||||||
self.conn = None
|
self.conn = None
|
||||||
self.initialize_db()
|
self.initialize_db()
|
||||||
|
|
||||||
def initialize_db(self):
|
def initialize_db(self):
|
||||||
"""Create the database tables if they don't exist."""
|
"""Create the database tables if they don't exist."""
|
||||||
try:
|
try:
|
||||||
self.conn = sqlite3.connect(self.db_path)
|
self.conn = sqlite3.connect(self.db_path)
|
||||||
cursor = self.conn.cursor()
|
cursor = self.conn.cursor()
|
||||||
|
|
||||||
# Create table for word performance
|
# Create table for word performance
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
CREATE TABLE IF NOT EXISTS word_performance (
|
CREATE TABLE IF NOT EXISTS word_performance (
|
||||||
id INTEGER PRIMARY KEY,
|
id INTEGER PRIMARY KEY,
|
||||||
word TEXT NOT NULL,
|
word TEXT NOT NULL,
|
||||||
|
|
@ -31,10 +34,12 @@ class WordsDatabase:
|
||||||
combined_score REAL NOT NULL,
|
combined_score REAL NOT NULL,
|
||||||
timestamp DATETIME DEFAULT CURRENT_TIMESTAMP
|
timestamp DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||||
)
|
)
|
||||||
''')
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
# Create table for word statistics (aggregated data)
|
# Create table for word statistics (aggregated data)
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
CREATE TABLE IF NOT EXISTS word_stats (
|
CREATE TABLE IF NOT EXISTS word_stats (
|
||||||
word TEXT PRIMARY KEY,
|
word TEXT PRIMARY KEY,
|
||||||
avg_improvement REAL NOT NULL,
|
avg_improvement REAL NOT NULL,
|
||||||
|
|
@ -45,31 +50,50 @@ class WordsDatabase:
|
||||||
best_position TEXT NOT NULL,
|
best_position TEXT NOT NULL,
|
||||||
last_updated DATETIME DEFAULT CURRENT_TIMESTAMP
|
last_updated DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||||
)
|
)
|
||||||
''')
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
self.conn.commit()
|
self.conn.commit()
|
||||||
print(f"Database initialized at {self.db_path}")
|
print(f"Database initialized at {self.db_path}")
|
||||||
except sqlite3.Error as e:
|
except sqlite3.Error as e:
|
||||||
print(f"Database error: {e}")
|
print(f"Database error: {e}")
|
||||||
|
|
||||||
def record_word_performance(self, word: str, position: str, benign_score: float,
|
def record_word_performance(
|
||||||
improvement: float, token_count: int, combined_score: float):
|
self,
|
||||||
|
word: str,
|
||||||
|
position: str,
|
||||||
|
benign_score: float,
|
||||||
|
improvement: float,
|
||||||
|
token_count: int,
|
||||||
|
combined_score: float,
|
||||||
|
):
|
||||||
"""Record the performance of a word when added to a prefix."""
|
"""Record the performance of a word when added to a prefix."""
|
||||||
if self.conn is None:
|
if self.conn is None:
|
||||||
self.initialize_db()
|
self.initialize_db()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
cursor = self.conn.cursor()
|
cursor = self.conn.cursor()
|
||||||
|
|
||||||
# Insert performance record
|
# Insert performance record
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
INSERT INTO word_performance
|
INSERT INTO word_performance
|
||||||
(word, position, benign_score, improvement, token_count, combined_score)
|
(word, position, benign_score, improvement, token_count, combined_score)
|
||||||
VALUES (?, ?, ?, ?, ?, ?)
|
VALUES (?, ?, ?, ?, ?, ?)
|
||||||
''', (word, position, benign_score, improvement, token_count, combined_score))
|
""",
|
||||||
|
(
|
||||||
|
word,
|
||||||
|
position,
|
||||||
|
benign_score,
|
||||||
|
improvement,
|
||||||
|
token_count,
|
||||||
|
combined_score,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
# Update statistics
|
# Update statistics
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
INSERT INTO word_stats
|
INSERT INTO word_stats
|
||||||
(word, avg_improvement, max_improvement, avg_token_count, min_token_count, use_count, best_position)
|
(word, avg_improvement, max_improvement, avg_token_count, min_token_count, use_count, best_position)
|
||||||
VALUES (?, ?, ?, ?, ?, 1, ?)
|
VALUES (?, ?, ?, ?, ?, 1, ?)
|
||||||
|
|
@ -81,104 +105,137 @@ class WordsDatabase:
|
||||||
use_count = use_count + 1,
|
use_count = use_count + 1,
|
||||||
best_position = CASE WHEN ? > max_improvement THEN ? ELSE best_position END,
|
best_position = CASE WHEN ? > max_improvement THEN ? ELSE best_position END,
|
||||||
last_updated = CURRENT_TIMESTAMP
|
last_updated = CURRENT_TIMESTAMP
|
||||||
''', (
|
""",
|
||||||
word, improvement, improvement, token_count, token_count, position,
|
(
|
||||||
improvement, improvement, token_count, token_count, improvement, position
|
word,
|
||||||
))
|
improvement,
|
||||||
|
improvement,
|
||||||
|
token_count,
|
||||||
|
token_count,
|
||||||
|
position,
|
||||||
|
improvement,
|
||||||
|
improvement,
|
||||||
|
token_count,
|
||||||
|
token_count,
|
||||||
|
improvement,
|
||||||
|
position,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
self.conn.commit()
|
self.conn.commit()
|
||||||
except sqlite3.Error as e:
|
except sqlite3.Error as e:
|
||||||
print(f"Error recording word performance: {e}")
|
print(f"Error recording word performance: {e}")
|
||||||
# Still try to continue without failing
|
# Still try to continue without failing
|
||||||
|
|
||||||
def get_top_words(self, limit: int = 20, min_uses: int = 2, sort_by: str = "improvement",
|
def get_top_words(
|
||||||
token_weight: float = 0.0) -> List[str]:
|
self,
|
||||||
|
limit: int = 20,
|
||||||
|
min_uses: int = 2,
|
||||||
|
sort_by: str = "improvement",
|
||||||
|
token_weight: float = 0.0,
|
||||||
|
) -> List[str]:
|
||||||
"""
|
"""
|
||||||
Get the top-performing words based on selected criteria.
|
Get the top-performing words based on selected criteria.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
-----------
|
-----------
|
||||||
limit: Maximum number of words to return
|
limit: Maximum number of words to return
|
||||||
min_uses: Minimum number of uses a word must have to be considered
|
min_uses: Minimum number of uses a word must have to be considered
|
||||||
sort_by: How to sort the results - options: "improvement", "tokens", "combined"
|
sort_by: How to sort the results - options: "improvement", "tokens", "combined"
|
||||||
token_weight: When sort_by="combined", weight for token count vs improvement (0-1)
|
token_weight: When sort_by="combined", weight for token count vs improvement (0-1)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
--------
|
--------
|
||||||
List of words matching the criteria
|
List of words matching the criteria
|
||||||
"""
|
"""
|
||||||
if self.conn is None:
|
if self.conn is None:
|
||||||
self.initialize_db()
|
self.initialize_db()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
cursor = self.conn.cursor()
|
cursor = self.conn.cursor()
|
||||||
|
|
||||||
# Different sorting strategies
|
# Different sorting strategies
|
||||||
if sort_by == "tokens":
|
if sort_by == "tokens":
|
||||||
# Sort by token count (ascending) then by improvement (descending)
|
# Sort by token count (ascending) then by improvement (descending)
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
SELECT word FROM word_stats
|
SELECT word FROM word_stats
|
||||||
WHERE use_count >= ? AND avg_improvement > 0
|
WHERE use_count >= ? AND avg_improvement > 0
|
||||||
ORDER BY min_token_count ASC, avg_improvement DESC
|
ORDER BY min_token_count ASC, avg_improvement DESC
|
||||||
LIMIT ?
|
LIMIT ?
|
||||||
''', (min_uses, limit))
|
""",
|
||||||
|
(min_uses, limit),
|
||||||
|
)
|
||||||
elif sort_by == "combined":
|
elif sort_by == "combined":
|
||||||
# Get all qualifying words with their stats
|
# Get all qualifying words with their stats
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
SELECT word, avg_improvement, min_token_count
|
SELECT word, avg_improvement, min_token_count
|
||||||
FROM word_stats
|
FROM word_stats
|
||||||
WHERE use_count >= ? AND avg_improvement > 0
|
WHERE use_count >= ? AND avg_improvement > 0
|
||||||
''', (min_uses,))
|
""",
|
||||||
|
(min_uses,),
|
||||||
|
)
|
||||||
|
|
||||||
# Calculate combined scores
|
# Calculate combined scores
|
||||||
results = cursor.fetchall()
|
results = cursor.fetchall()
|
||||||
if not results:
|
if not results:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
# Normalize values
|
# Normalize values
|
||||||
max_improvement = max(row[1] for row in results)
|
max_improvement = max(row[1] for row in results)
|
||||||
max_tokens = max(row[2] for row in results)
|
max_tokens = max(row[2] for row in results)
|
||||||
|
|
||||||
# Calculate combined score for each word
|
# Calculate combined score for each word
|
||||||
scored_words = []
|
scored_words = []
|
||||||
for row in results:
|
for row in results:
|
||||||
word = row[0]
|
word = row[0]
|
||||||
norm_improvement = row[1] / max_improvement if max_improvement > 0 else 0
|
norm_improvement = row[1] / max_improvement if max_improvement > 0 else 0
|
||||||
norm_tokens = 1 - (row[2] / max_tokens if max_tokens > 0 else 0) # Invert so lower is better
|
norm_tokens = 1 - (
|
||||||
combined_score = (1 - token_weight) * norm_improvement + token_weight * norm_tokens
|
row[2] / max_tokens if max_tokens > 0 else 0
|
||||||
|
) # Invert so lower is better
|
||||||
|
combined_score = (
|
||||||
|
1 - token_weight
|
||||||
|
) * norm_improvement + token_weight * norm_tokens
|
||||||
scored_words.append((word, combined_score))
|
scored_words.append((word, combined_score))
|
||||||
|
|
||||||
# Sort by combined score and return top words
|
# Sort by combined score and return top words
|
||||||
scored_words.sort(key=lambda x: x[1], reverse=True)
|
scored_words.sort(key=lambda x: x[1], reverse=True)
|
||||||
return [word for word, _ in scored_words[:limit]]
|
return [word for word, _ in scored_words[:limit]]
|
||||||
else:
|
else:
|
||||||
# Default: sort by improvement
|
# Default: sort by improvement
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
SELECT word FROM word_stats
|
SELECT word FROM word_stats
|
||||||
WHERE use_count >= ? AND avg_improvement > 0
|
WHERE use_count >= ? AND avg_improvement > 0
|
||||||
ORDER BY avg_improvement DESC
|
ORDER BY avg_improvement DESC
|
||||||
LIMIT ?
|
LIMIT ?
|
||||||
''', (min_uses, limit))
|
""",
|
||||||
|
(min_uses, limit),
|
||||||
|
)
|
||||||
|
|
||||||
results = cursor.fetchall()
|
results = cursor.fetchall()
|
||||||
return [row[0] for row in results]
|
return [row[0] for row in results]
|
||||||
except sqlite3.Error as e:
|
except sqlite3.Error as e:
|
||||||
print(f"Error getting top words: {e}")
|
print(f"Error getting top words: {e}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def get_word_stats(self, word: str) -> Optional[Dict[str, Any]]:
|
def get_word_stats(self, word: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Get statistics for a specific word."""
|
"""Get statistics for a specific word."""
|
||||||
if self.conn is None:
|
if self.conn is None:
|
||||||
self.initialize_db()
|
self.initialize_db()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
cursor = self.conn.cursor()
|
cursor = self.conn.cursor()
|
||||||
cursor.execute('''
|
cursor.execute(
|
||||||
|
"""
|
||||||
SELECT word, avg_improvement, max_improvement, avg_token_count, min_token_count, use_count, best_position
|
SELECT word, avg_improvement, max_improvement, avg_token_count, min_token_count, use_count, best_position
|
||||||
FROM word_stats
|
FROM word_stats
|
||||||
WHERE word = ?
|
WHERE word = ?
|
||||||
''', (word,))
|
""",
|
||||||
|
(word,),
|
||||||
|
)
|
||||||
|
|
||||||
result = cursor.fetchone()
|
result = cursor.fetchone()
|
||||||
if result:
|
if result:
|
||||||
return {
|
return {
|
||||||
|
|
@ -188,15 +245,15 @@ class WordsDatabase:
|
||||||
"avg_token_count": result[3],
|
"avg_token_count": result[3],
|
||||||
"min_token_count": result[4],
|
"min_token_count": result[4],
|
||||||
"use_count": result[5],
|
"use_count": result[5],
|
||||||
"best_position": result[6]
|
"best_position": result[6],
|
||||||
}
|
}
|
||||||
return None
|
return None
|
||||||
except sqlite3.Error as e:
|
except sqlite3.Error as e:
|
||||||
print(f"Error getting word stats: {e}")
|
print(f"Error getting word stats: {e}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
"""Close the database connection."""
|
"""Close the database connection."""
|
||||||
if self.conn:
|
if self.conn:
|
||||||
self.conn.close()
|
self.conn.close()
|
||||||
self.conn = None
|
self.conn = None
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue