Use black to format

This commit is contained in:
Joey Yakimowich-Payne 2025-05-12 17:47:18 -06:00
commit 06eac3eaba
4 changed files with 563 additions and 794 deletions

View file

@ -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
View file

@ -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
View file

@ -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]))

View file

@ -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