| |
| """Interactive inference example for authorized red-team evaluation. |
| |
| This script generates attacker queries but deliberately does not call a target |
| model endpoint. A human operator must paste each authorized target response. |
| """ |
|
|
| import argparse |
| import json |
| from pathlib import Path |
| from typing import Dict, List, Optional |
|
|
| import torch |
| from huggingface_hub import hf_hub_download |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
|
|
|
|
| Message = Dict[str, str] |
|
|
|
|
| def resolve_prompt_config(model: str, explicit_path: Optional[Path]) -> Path: |
| """Resolve model-specific prompt settings locally or from the Hub.""" |
| if explicit_path is not None: |
| return explicit_path |
|
|
| local_candidate = Path(model) / "prompt_template.json" |
| if local_candidate.is_file(): |
| return local_candidate |
|
|
| return Path( |
| hf_hub_download(repo_id=model, filename="prompt_template.json") |
| ) |
|
|
|
|
| def load_prompt_config(path: Path) -> Dict: |
| with path.open("r", encoding="utf-8") as handle: |
| prompt_config = json.load(handle) |
|
|
| required_keys = { |
| "system_prompt", |
| "initial_user_prompt_template", |
| "target_response_role", |
| "attacker_response_role", |
| "max_interaction_turns", |
| } |
| missing_keys = sorted(required_keys - prompt_config.keys()) |
| if missing_keys: |
| raise ValueError( |
| f"Prompt config {path} is missing keys: {', '.join(missing_keys)}" |
| ) |
| if prompt_config["target_response_role"] != "user": |
| raise ValueError("The TRACE target-response role must be 'user'.") |
| if prompt_config["attacker_response_role"] != "assistant": |
| raise ValueError("The TRACE attacker-response role must be 'assistant'.") |
| return prompt_config |
|
|
|
|
| def build_initial_messages(prompt_config: Dict, objective: str) -> List[Message]: |
| initial_prompt = prompt_config["initial_user_prompt_template"].format( |
| harmful_target=objective |
| ) |
| return [ |
| {"role": "system", "content": prompt_config["system_prompt"]}, |
| {"role": "user", "content": initial_prompt}, |
| ] |
|
|
|
|
| def append_completed_turn( |
| messages: List[Message], attacker_query: str, target_response: str |
| ) -> None: |
| messages.extend( |
| [ |
| {"role": "assistant", "content": attacker_query}, |
| |
| |
| {"role": "user", "content": target_response}, |
| ] |
| ) |
|
|
|
|
| def generate_attacker_query( |
| model, |
| tokenizer, |
| messages: List[Message], |
| max_new_tokens: int, |
| do_sample: bool, |
| temperature: float, |
| top_p: float, |
| top_k: int, |
| ) -> str: |
| model_inputs = tokenizer.apply_chat_template( |
| messages, |
| tokenize=True, |
| add_generation_prompt=True, |
| return_tensors="pt", |
| return_dict=True, |
| ) |
| model_inputs = model_inputs.to(model.device) |
|
|
| generation_kwargs = { |
| "max_new_tokens": max_new_tokens, |
| "do_sample": do_sample, |
| "pad_token_id": tokenizer.pad_token_id, |
| "eos_token_id": tokenizer.eos_token_id, |
| } |
| if do_sample: |
| generation_kwargs.update( |
| {"temperature": temperature, "top_p": top_p, "top_k": top_k} |
| ) |
|
|
| with torch.inference_mode(): |
| output_ids = model.generate(**model_inputs, **generation_kwargs) |
|
|
| generated_ids = output_ids[0, model_inputs.input_ids.shape[1] :] |
| return tokenizer.decode(generated_ids, skip_special_tokens=True).strip() |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser( |
| description="Interactive TRACE attacker inference for authorized evaluation." |
| ) |
| parser.add_argument("--model", required=True, help="Local path or Hub model ID.") |
| parser.add_argument( |
| "--objective", |
| required=True, |
| help="Authorized red-team objective supplied to the attacker policy.", |
| ) |
| parser.add_argument( |
| "--prompt-config", |
| type=Path, |
| default=None, |
| help=( |
| "Optional path to prompt_template.json. By default it is loaded " |
| "from the local model directory or the model's Hub repository." |
| ), |
| ) |
| parser.add_argument("--max-turns", type=int, default=None) |
| parser.add_argument("--max-new-tokens", type=int, default=128) |
| parser.add_argument( |
| "--do-sample", |
| action=argparse.BooleanOptionalAction, |
| default=True, |
| help="Sample attacker outputs; enabled in the reported validation setting.", |
| ) |
| parser.add_argument("--temperature", type=float, default=0.5) |
| parser.add_argument("--top-p", type=float, default=0.9) |
| parser.add_argument( |
| "--top-k", |
| type=int, |
| default=0, |
| help=( |
| "Transformers top-k value. Zero disables top-k and corresponds " |
| "to top_k=-1 in the training-time vLLM configuration." |
| ), |
| ) |
| parser.add_argument( |
| "--device-map", |
| default="auto", |
| help="Transformers device_map value; defaults to auto.", |
| ) |
| return parser.parse_args() |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
| prompt_config_path = resolve_prompt_config(args.model, args.prompt_config) |
| prompt_config = load_prompt_config(prompt_config_path) |
| max_turns = ( |
| int(prompt_config["max_interaction_turns"]) |
| if args.max_turns is None |
| else args.max_turns |
| ) |
| if max_turns <= 0: |
| raise ValueError("--max-turns must be positive") |
| if args.max_new_tokens <= 0: |
| raise ValueError("--max-new-tokens must be positive") |
| if args.do_sample and args.temperature <= 0: |
| raise ValueError("--temperature must be positive when sampling is enabled") |
| if not 0 < args.top_p <= 1: |
| raise ValueError("--top-p must be in the interval (0, 1]") |
| if args.top_k < 0: |
| raise ValueError("--top-k must be non-negative for Transformers") |
|
|
| tokenizer = AutoTokenizer.from_pretrained(args.model) |
| model = AutoModelForCausalLM.from_pretrained( |
| args.model, |
| torch_dtype="auto", |
| device_map=args.device_map, |
| ).eval() |
|
|
| messages = build_initial_messages(prompt_config, args.objective) |
| for turn_index in range(1, max_turns + 1): |
| attacker_query = generate_attacker_query( |
| model=model, |
| tokenizer=tokenizer, |
| messages=messages, |
| max_new_tokens=args.max_new_tokens, |
| do_sample=args.do_sample, |
| temperature=args.temperature, |
| top_p=args.top_p, |
| top_k=args.top_k, |
| ) |
| print(f"\n[attacker turn {turn_index}]\n{attacker_query}\n", flush=True) |
|
|
| if turn_index == max_turns: |
| break |
| try: |
| target_response = input( |
| "Paste the authorized target-model response " |
| "(or type /stop to finish):\n" |
| ).strip() |
| except EOFError: |
| break |
| if target_response == "/stop": |
| break |
| if not target_response: |
| print("Empty target response; stopping without adding an invalid turn.") |
| break |
| append_completed_turn(messages, attacker_query, target_response) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|