mirror of
https://github.com/msoedov/agentic_security.git
synced 2026-09-30 19:59:34 +02:00
Merge pull request #224 from Mundi-Xu/datasets-optimize
refactor: standardize CSV loading from ./datasets and improve robustness
This commit is contained in:
1 file changed
+13
-7
@@ -248,13 +248,13 @@ def load_jailbreak_v28k() -> ProbeDataset:
|
|||||||
@cache_to_disk()
|
@cache_to_disk()
|
||||||
def load_local_csv() -> ProbeDataset:
|
def load_local_csv() -> ProbeDataset:
|
||||||
"""Load prompts from local CSV files."""
|
"""Load prompts from local CSV files."""
|
||||||
csv_files = [f for f in os.listdir(".") if f.endswith(".csv")]
|
csv_files = [f for f in os.listdir("./datasets") if f.endswith(".csv")]
|
||||||
logger.info(f"Found {len(csv_files)} CSV files: {csv_files}")
|
logger.info(f"Found {len(csv_files)} CSV files: {csv_files}")
|
||||||
|
|
||||||
prompts = []
|
prompts = []
|
||||||
for file in csv_files:
|
for file in csv_files:
|
||||||
try:
|
try:
|
||||||
df = pd.read_csv(file)
|
df = pd.read_csv(os.path.join("./datasets", file), encoding_errors="ignore")
|
||||||
if "prompt" in df.columns:
|
if "prompt" in df.columns:
|
||||||
prompts.extend(df["prompt"].tolist())
|
prompts.extend(df["prompt"].tolist())
|
||||||
else:
|
else:
|
||||||
@@ -270,7 +270,7 @@ def load_csv(file: str) -> ProbeDataset:
|
|||||||
"""Load prompts from local CSV files."""
|
"""Load prompts from local CSV files."""
|
||||||
prompts = []
|
prompts = []
|
||||||
try:
|
try:
|
||||||
df = pd.read_csv(file)
|
df = pd.read_csv(os.path.join("./datasets", file), encoding_errors="ignore")
|
||||||
prompts = df["prompt"].tolist()
|
prompts = df["prompt"].tolist()
|
||||||
if "prompt" in df.columns:
|
if "prompt" in df.columns:
|
||||||
prompts.extend(df["prompt"].tolist())
|
prompts.extend(df["prompt"].tolist())
|
||||||
@@ -284,14 +284,14 @@ def load_csv(file: str) -> ProbeDataset:
|
|||||||
@cache_to_disk(1)
|
@cache_to_disk(1)
|
||||||
def load_local_csv_files() -> list[ProbeDataset]:
|
def load_local_csv_files() -> list[ProbeDataset]:
|
||||||
"""Load prompts from local CSV files and return a list of ProbeDataset objects."""
|
"""Load prompts from local CSV files and return a list of ProbeDataset objects."""
|
||||||
csv_files = [f for f in os.listdir(".") if f.endswith(".csv")]
|
csv_files = [f for f in os.listdir("./datasets") if f.endswith(".csv")]
|
||||||
logger.info(f"Found {len(csv_files)} CSV files: {csv_files}")
|
logger.info(f"Found {len(csv_files)} CSV files: {csv_files}")
|
||||||
|
|
||||||
datasets = []
|
datasets = []
|
||||||
|
|
||||||
for file in csv_files:
|
for file in csv_files:
|
||||||
try:
|
try:
|
||||||
df = pd.read_csv(file)
|
df = pd.read_csv(os.path.join("./datasets", file), encoding_errors="ignore")
|
||||||
if "prompt" in df.columns:
|
if "prompt" in df.columns:
|
||||||
prompts = df["prompt"].tolist()
|
prompts = df["prompt"].tolist()
|
||||||
datasets.append(create_probe_dataset(file, prompts, {"src": file}))
|
datasets.append(create_probe_dataset(file, prompts, {"src": file}))
|
||||||
@@ -336,12 +336,18 @@ class StenographyTransformer:
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
size = len(dataset.prompts)
|
# Convert async generator (if any) to a list
|
||||||
|
prompts = list(dataset.prompts) if hasattr(dataset.prompts, '__iter__') else []
|
||||||
|
size = len(prompts)
|
||||||
|
|
||||||
|
if size == 0:
|
||||||
|
continue
|
||||||
|
|
||||||
sample_size = max(1, int(size * 0.25)) # At least 1 sample
|
sample_size = max(1, int(size * 0.25)) # At least 1 sample
|
||||||
|
|
||||||
for name, fn in self.TRANSFORMATIONS.items():
|
for name, fn in self.TRANSFORMATIONS.items():
|
||||||
logger.info(f"Applying {name} to {dataset.dataset_name}")
|
logger.info(f"Applying {name} to {dataset.dataset_name}")
|
||||||
sampled = random.sample(dataset.prompts, min(sample_size, size))
|
sampled = random.sample(prompts, min(sample_size, size))
|
||||||
try:
|
try:
|
||||||
transformed = [
|
transformed = [
|
||||||
f"Follow the instructions in {name}: {fn(str(p))}"
|
f"Follow the instructions in {name}: {fn(str(p))}"
|
||||||
|
|||||||
Reference in new issue
Block a user