MCPcopy Create free account
hub / github.com/SalesforceAIResearch/perfcodegen / Evaluator

Class Evaluator

src/evaluate.py:103–665  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

101
102
103class Evaluator(object):
104 def __init__(self, dataset, dataset_repo = "test_datasets", mbpp_helpfile = None):
105 self.dataset = Dataset(dataset, data_path = os.path.join(dataset_repo, dataset.replace("/", "_")), testfile_path = mbpp_helpfile)
106 self.dataset.load_testcases()
107 self.dataset.load_groundtruths()
108
109
110 def load_solutions(self, solution_file):
111 self.solution_file = solution_file
112 self.solutions = json.load(open(self.solution_file, "r"))
113
114 def check_element_type(lst, t):
115 for l in lst:
116 if not isinstance(l, t):
117 if t == float and isinstance(l, int):
118 continue
119 return False
120
121 return True
122
123 def transform_element_type(lst, t):
124 new_lst = []
125 for l in lst:
126 new_lst.append(t(l))
127 return new_lst
128
129 def prepare_testcases(self, solution_io, gt_io, testcases):
130 if solution_io == gt_io:
131 return testcases
132 elif solution_io:
133 new_testcases = []
134 for testcase in testcases:
135 new_testcase = {}
136 new_testcase["output"] = testcase["output"]
137 if check_element_type(testcase["input"], str):
138 new_testcase["input"] = ["\n".join(testcase["input"]) + "\n"]
139 elif check_element_type(testcase["input"], float):
140 new_testcase["input"] = ["\n".join(self.transform_element_type(testcase["input"], str)) + "\n"]
141 else:
142 new_testcase["input"] = testcase["input"]
143 new_testcases.append(new_testcase)
144 return new_testcases
145 else:
146 new_testcases = []
147 for testcase in testcases:
148 new_testcase = {}
149 new_testcase["output"] = testcase["output"]
150 if len(testcase["input"]) == 1 and isinstance(testcase["input"][0], str):
151 new_testcase["input"] = [testcase["input"][0].split("\n")]
152 else:
153 new_testcase["input"] = testcase["input"]
154 new_testcases.append(new_testcase)
155 return new_testcases
156
157
158 def get_expected_outputs(self):
159 prompts, overlong_prompts = self.dataset.get_all_prompts()
160

Callers 2

run_correctness_checkMethod · 0.90
run_time_measurementMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected