From 25b12a512d2f833b64ae6d1312cf475f7fb47b59 Mon Sep 17 00:00:00 2001 From: Joey Yakimowich-Payne Date: Tue, 13 May 2025 14:58:44 -0600 Subject: [PATCH] Update install stuff --- hacking.py | 3 +-- installnotes.bash | 13 +++---------- utils.py | 17 +++++++++++++++++ 3 files changed, 21 insertions(+), 12 deletions(-) diff --git a/hacking.py b/hacking.py index 54f2b9d..9ca7459 100644 --- a/hacking.py +++ b/hacking.py @@ -8,17 +8,16 @@ import argparse from typing import List, Dict, Tuple, Any, Optional from transformers import pipeline, AutoTokenizer, AutoModelForSequenceClassification from huggingface_hub import login -from llm_attacks.minimal_gcg.opt_utils import get_filtered_cands 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, + get_filtered_cands, ) # check if cuda is available diff --git a/installnotes.bash b/installnotes.bash index e8bf9e5..62bc924 100644 --- a/installnotes.bash +++ b/installnotes.bash @@ -3,22 +3,15 @@ conda init conda create --name prompt-guard python=3.12 -c conda-forge conda activate prompt-guard -curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -rustup install 1.82 -rustup default 1.82 - -pip install fschat #cuda -pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128 +pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu126 #or cpu pip install torch --index-url https://download.pytorch.org/whl/cpu pip install -U "huggingface_hub[cli]" # install transformers -pip install transformers -export RUSTFLAGS="-A invalid_reference_casting" -pip install git+https://github.com/llm-attacks/llm-attacks +pip install -U transformers +pip install -U tiktoken huggingface-cli login -pip install -U transformers diff --git a/utils.py b/utils.py index 4be7959..a3ea921 100644 --- a/utils.py +++ b/utils.py @@ -644,3 +644,20 @@ def get_combined_score( best_idx = idx return best_idx + +def get_filtered_cands(tokenizer, control_cand, filter_cand=True, curr_control=None): + cands, count = [], 0 + for i in range(control_cand.shape[0]): + decoded_str = tokenizer.decode(control_cand[i], skip_special_tokens=True) + if filter_cand: + if decoded_str != curr_control and len(tokenizer(decoded_str, add_special_tokens=False).input_ids) == len(control_cand[i]): + cands.append(decoded_str) + else: + count += 1 + else: + cands.append(decoded_str) + + if filter_cand: + cands = cands + [cands[-1]] * (len(control_cand) - len(cands)) + # print(f"Warning: {round(count / len(control_cand), 2)} control candidates were not valid") + return cands \ No newline at end of file