CoolFace
Datasetpublic

PerturbReason/PerturbReason_dataset_code

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes12downloads
qwen_rescue_vllm.py96 linesDownload Raw Back to eval_v3
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