| 84 | json.dump(record, f, ensure_ascii=False) |
| 85 | f.write('\n') |
| 86 | def solve_init(args): |
| 87 | if DEBUG: |
| 88 | ckpt = LLM( |
| 89 | model=args.model_path, |
| 90 | tensor_parallel_size=1 |
| 91 | ) |
| 92 | else: |
| 93 | ckpt = LLM( |
| 94 | model=args.model_path, |
| 95 | tensor_parallel_size=args.tensor_parallel_size |
| 96 | ) |
| 97 | print("ckpt is ready.") |
| 98 | # 读取 JSON 数据 |
| 99 | dataset = load_dataset('json', data_files=args.dev_dataset_path)['train'] |
| 100 | # 如果处于调试模式,只取少量数据 |
| 101 | if DEBUG: |
| 102 | # 获取数据集大小 |
| 103 | dataset_size = len(dataset) |
| 104 | # 随机采样8个索引(或小于数据集大小的数量) |
| 105 | sample_size = min(8, dataset_size) |
| 106 | sampled_indices = random.sample(range(dataset_size), sample_size) |
| 107 | # 只保留采样的数据 |
| 108 | dataset = dataset.select(sampled_indices) |
| 109 | records = [] |
| 110 | for i, data in enumerate(dataset): |
| 111 | record = {} |
| 112 | record['question'] = data['question'] |
| 113 | record['golden_answers'] = data['golden_answers'] |
| 114 | record['state'] = "undo" |
| 115 | record['resample_times'] = 0 |
| 116 | records.append(record) |
| 117 | return ckpt , records |
| 118 | def generate_naive_generation_cot_prompt(question): |
| 119 | system_message = """You are a helpful assistant that thinks through problems step by step before providing a final answer based on your own knowledge. |
| 120 | |