PerturbReason/PerturbReason_dataset_code
012
1#!/usr/bin/env python32"""3qwen_rescue_vllm.py4====================5Batch vLLM inference for Tier-4 LLM rescue prompts.6Same pattern as qwen_batch_vllm.py — reads JSONL, runs batch inference, writes JSONL.7 8Usage::9 10 python qwen_rescue_vllm.py \11 --model_path /path/to/Qwen3-4B-Instruct \12 --input_file eval_v3/eval_v3_output/rescue_prompts.jsonl \13 --output_file eval_v3/eval_v3_output/rescue_responses.jsonl \14 --max_new_tokens 25615"""16 17import argparse18import json19from vllm import LLM, SamplingParams20 21 22def parse_args():23 parser = argparse.ArgumentParser(24 description="Batch vLLM inference for eval_v3 Tier-4 rescue")25 parser.add_argument("--model_path", type=str, required=True)26 parser.add_argument("--input_file", type=str, required=True,27 help="Rescue prompts JSONL (from export_rescue.py)")28 parser.add_argument("--output_file", type=str, required=True,29 help="Output JSONL with rescue responses")30 parser.add_argument("--max_new_tokens", type=int, default=256)31 parser.add_argument("--tp_size", type=int, default=1,32 help="Tensor parallel size")33 parser.add_argument("--temperature", type=float, default=0.1,34 help="Low temp for deterministic judging")35 parser.add_argument("--top_p", type=float, default=0.9)36 return parser.parse_args()37 38 39def main():40 args = parse_args()41 42 # 1. Load model43 print(f"Loading model from {args.model_path} ...")44 llm = LLM(45 model=args.model_path,46 tensor_parallel_size=args.tp_size,47 trust_remote_code=True,48 dtype="bfloat16",49 max_model_len=8192,50 )51 52 sampling_params = SamplingParams(53 temperature=args.temperature,54 top_p=args.top_p,55 max_tokens=args.max_new_tokens,56 )57 58 # 2. Load rescue prompts59 records = []60 prompts = []61 with open(args.input_file, "r", encoding="utf-8") as f:62 for line in f:63 line = line.strip()64 if not line:65 continue66 record = json.loads(line)67 records.append(record)68 prompts.append(record["prompt"])69 70 print(f"Loaded {len(prompts)} rescue prompts from {args.input_file}")71 72 if not prompts:73 print("No prompts to process.")74 return75 76 # 3. Batch inference77 print("Running batch inference ...")78 outputs = llm.generate(prompts, sampling_params)79 80 # 4. Write results81 with open(args.output_file, "w", encoding="utf-8") as f_out:82 for record, output in zip(records, outputs):83 response_text = output.outputs[0].text84 result_entry = {85 "prompt": record["prompt"],86 "metadata": record.get("metadata", {}),87 "rescue_response": response_text,88 }89 f_out.write(json.dumps(result_entry, ensure_ascii=False) + "\n")90 91 print(f"Wrote {len(records)} rescue responses → {args.output_file}")92 93 94if __name__ == "__main__":95 main()96 