CoolFace
Datasetpublic

PerturbReason/PerturbReason_dataset_code

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes12downloads
qwen_batch_vllm_0419.py77 linesDownload Raw Back to Baselines
1import argparse2import json3import os4import glob5from vllm import LLM, SamplingParams6 7def parse_args():8    parser = argparse.ArgumentParser()9    parser.add_argument("--model_path", type=str, required=True)10    parser.add_argument("--input_dir", type=str, required=True)11    parser.add_argument("--output_dir", type=str, required=True)12    parser.add_argument("--max_new_tokens", type=int, default=4096)13    parser.add_argument("--tp_size", type=int, default=2)14    return parser.parse_args()15 16def main():17    args = parse_args()18 19    print(f"Loading Qwen via vLLM from {args.model_path}...")20    llm = LLM(21        model=args.model_path,22        tensor_parallel_size=args.tp_size,23        trust_remote_code=True,24        dtype="bfloat16",25        max_model_len=3276826    )27 28    sampling_params = SamplingParams(29        temperature=0.7,30        top_p=0.9,31        max_tokens=args.max_new_tokens32    )33 34    os.makedirs(args.output_dir, exist_ok=True)35    input_files = glob.glob(os.path.join(args.input_dir, "*.jsonl"))36    print(f"Found {len(input_files)} JSONL files in {args.input_dir}")37 38    for file_path in input_files:39        file_name = os.path.basename(file_path)40        output_path = os.path.join(args.output_dir, f"qwen_pred_{file_name}")41        print(f"Processing: {file_name} -> {output_path}")42 43        prompts = []44        original_records = []45        with open(file_path, 'r', encoding='utf-8') as f:46            for line in f:47                if not line.strip():48                    continue49                try:50                    record = json.loads(line)51                    if record.get("prompt"):52                        prompts.append(record["prompt"])53                        original_records.append(record)54                except json.JSONDecodeError:55                    continue56 57        if not prompts:58            continue59 60        outputs = llm.generate(prompts, sampling_params)61 62        with open(output_path, 'w', encoding='utf-8') as f_out:63            for record, output in zip(original_records, outputs):64                generated_text = output.outputs[0].text65                result_entry = {66                    "source_file": file_name,67                    "prompt": record["prompt"],68                    "ground_truth_response": record.get("response", ""),69                    "model_output": generated_text,70                    "path_ambiguity": record.get("path_ambiguity", ""),71                    "data_type": record.get("data_type", ""),72                }73                f_out.write(json.dumps(result_entry) + "\n")74 75if __name__ == "__main__":76    main()77