CoolFace
Apppublic

KChad/Prompt-Injection-RL-environment

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
curate_llmail_dataset.py79 linesDownload Raw Back to scripts
1from __future__ import annotations2 3import argparse4import sys5from pathlib import Path6 7ROOT = Path(__file__).resolve().parents[1]8sys.path.insert(0, str(ROOT))9 10from env.llmail_curator import (11    DIFFICULTIES,12    curate_rows,13    load_llmail_contexts,14    load_rows_from_hf,15    load_rows_from_jsonl,16    write_curated_scenarios,17)18 19DEFAULT_OUTPUT_DIR = ROOT / "env" / "data" / "scenarios" / "curated"20 21 22def parse_args() -> argparse.Namespace:23    parser = argparse.ArgumentParser(24        description="Curate LLMail submissions into local easy/medium/hard scenario files for the IDPI mail benchmark."25    )26    parser.add_argument("--phase", action="append", choices=["Phase1", "Phase2"], help="Repeat to choose one or both Hugging Face splits. Defaults to both.")27    parser.add_argument("--input-jsonl-phase1", type=Path, help="Optional local Phase1 JSONL file.")28    parser.add_argument("--input-jsonl-phase2", type=Path, help="Optional local Phase2 JSONL file.")29    parser.add_argument("--scenarios-json", type=Path, help="Optional local copy of LLMail scenarios.json.")30    parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)31    parser.add_argument("--prefix", default="")32    parser.add_argument("--max-per-difficulty", type=int, default=50)33    parser.add_argument("--sample-limit", type=int, default=None, help="Optional cap per split for quick dry runs.")34    parser.add_argument("--seed", type=int, default=42)35    parser.add_argument("--overwrite", action="store_true")36    return parser.parse_args()37 38 39def rows_for_phase(phase: str, args: argparse.Namespace):40    if phase == "Phase1" and args.input_jsonl_phase1:41        return load_rows_from_jsonl(args.input_jsonl_phase1, sample_limit=args.sample_limit)42    if phase == "Phase2" and args.input_jsonl_phase2:43        return load_rows_from_jsonl(args.input_jsonl_phase2, sample_limit=args.sample_limit)44    return load_rows_from_hf(phase, sample_limit=args.sample_limit)45 46 47def main() -> None:48    args = parse_args()49    phases = args.phase or ["Phase1", "Phase2"]50    contexts = load_llmail_contexts(local_path=args.scenarios_json)51 52    merged = {difficulty: [] for difficulty in DIFFICULTIES}53    for phase in phases:54        curated = curate_rows(55            rows_for_phase(phase, args),56            phase=phase,57            contexts=contexts,58            max_per_difficulty=args.max_per_difficulty,59            seed=args.seed,60        )61        for difficulty in DIFFICULTIES:62            merged[difficulty].extend(curated[difficulty])63 64    written = write_curated_scenarios(65        merged,66        output_dir=args.output_dir,67        prefix=args.prefix,68        overwrite=args.overwrite,69    )70 71    for path in written:72        with path.open("r", encoding="utf-8-sig") as handle:73            count = sum(1 for _ in handle)74        print(f"wrote {count:>3} scenarios -> {path}")75 76 77if __name__ == "__main__":78    main()79