(data_name, prompt_type, samples: list=None, file_path: str=None, max_num_samples=None, execute=False)
| 51 | |
| 52 | |
| 53 | def evaluate(data_name, prompt_type, samples: list=None, file_path: str=None, max_num_samples=None, execute=False): |
| 54 | assert samples or file_path, "samples or file_path must be provided" |
| 55 | if not samples: |
| 56 | samples = list(load_jsonl(file_path)) |
| 57 | if 'idx' in samples[0]: |
| 58 | samples = {sample['idx']: sample for sample in samples}.values() |
| 59 | samples = sorted(samples, key=lambda x: x['idx']) |
| 60 | else: |
| 61 | samples = [dict(idx=idx, **sample) for idx, sample in enumerate(samples)] |
| 62 | |
| 63 | if max_num_samples: |
| 64 | print(f"max_num_samples: {max_num_samples} / {len(samples)}") |
| 65 | samples = samples[:max_num_samples] |
| 66 | |
| 67 | # parse gt |
| 68 | for sample in samples: |
| 69 | sample['gt_cot'], sample['gt'] = parse_ground_truth(sample, data_name) |
| 70 | params = [(idx, pred, sample['gt']) for idx, sample in enumerate(samples) for pred in sample['pred']] |
| 71 | |
| 72 | scores = [] |
| 73 | timeout_cnt = 0 |
| 74 | |
| 75 | with ProcessPool(max_workers=1) as pool: |
| 76 | future = pool.map(new_math_equal_process, params, timeout=3) |
| 77 | iterator = future.result() |
| 78 | with tqdm(total=len(samples), desc="Evaluate") as progress_bar: |
| 79 | while True: |
| 80 | try: |
| 81 | result = next(iterator) |
| 82 | scores.append(result) |
| 83 | except StopIteration: |
| 84 | break |
| 85 | except TimeoutError as error: |
| 86 | print(error) |
| 87 | scores.append(False) |
| 88 | timeout_cnt += 1 |
| 89 | except Exception as error: |
| 90 | print(error.traceback) |
| 91 | exit() |
| 92 | progress_bar.update(1) |
| 93 | # for debug only |
| 94 | # import random |
| 95 | # scores = [random.random() > 0.9 for _ in range(len(params))] |
| 96 | |
| 97 | idx = 0 |
| 98 | score_mat = [] |
| 99 | for sample in samples: |
| 100 | sample['score'] = scores[idx: idx+len(sample['pred'])] |
| 101 | assert len(sample['score']) == len(sample['pred']) |
| 102 | score_mat.append(sample['score']) |
| 103 | idx += len(sample['pred']) |
| 104 | |
| 105 | max_len = max([len(s) for s in score_mat]) |
| 106 | |
| 107 | for i, s in enumerate(score_mat): |
| 108 | if len(s) < max_len: |
| 109 | score_mat[i] = s + [s[-1]] * (max_len - len(s)) # pad |
| 110 |
no test coverage detected