(self, recorder, samples)
| 217 | ) |
| 218 | |
| 219 | def eval_sample_batch(self, recorder, samples): |
| 220 | id = 0 |
| 221 | # for sample in samples: |
| 222 | data, ideal = [], [] |
| 223 | for i in range(len(samples)): |
| 224 | prompt, correct_answer = self.pre_process(samples[i]) |
| 225 | data.extend([prompt] * self.num_samples_per_task) |
| 226 | ideal.append(correct_answer) |
| 227 | |
| 228 | response = self.completion_fn( |
| 229 | inputs=data, |
| 230 | temperature=self.temperature, |
| 231 | do_sample=True, |
| 232 | top_p=0.95, |
| 233 | max_tokens=self.max_tokens, |
| 234 | ) |
| 235 | for id in range(len(samples)): |
| 236 | results = response[ |
| 237 | id * self.num_samples_per_task : (id + 1) * self.num_samples_per_task |
| 238 | ] |
| 239 | prompt = data[id * self.num_samples_per_task] |
| 240 | correct_answer = ideal[id] |
| 241 | |
| 242 | stop_sequences = ( |
| 243 | ["\nclass", "\ndef", "\n#", "\nif"] |
| 244 | if self.dataset == "humaneval" |
| 245 | else ["\n[DONE]"] |
| 246 | ) |
| 247 | for i in range(len(results)): |
| 248 | for x in stop_sequences: |
| 249 | if x in results[i]: |
| 250 | results[i] = results[i].split(x)[0] |
| 251 | |
| 252 | if self.dataset == "humaneval": |
| 253 | results = [ |
| 254 | prompt[len("Complete the code:\n") :] + item for item in results |
| 255 | ] |
| 256 | |
| 257 | for i in range(len(results)): |
| 258 | with recorder.as_default_recorder(id): |
| 259 | record_sampling(prompt=prompt, sampled=results[i]) |
| 260 | |
| 261 | sampled = [["a" if item == None else item for item in results]] |
| 262 | pass_at_k, results = compute( |
| 263 | references=correct_answer, predictions=sampled, k=self.k |
| 264 | ) |
| 265 | with recorder.as_default_recorder(id): |
| 266 | if self.dataset == "humaneval": |
| 267 | evals.record.record_metrics( |
| 268 | pass_at_1=pass_at_k["pass@1"], |
| 269 | pass_at_10=pass_at_k["pass@10"], |
| 270 | pass_at_100=pass_at_k["pass@100"], |
| 271 | ) |
| 272 | elif self.dataset == "mbpp": |
| 273 | evals.record.record_metrics( |
| 274 | pass_at_1=pass_at_k["pass@1"], |
| 275 | pass_at_80=pass_at_k["pass@80"], |
| 276 | ) |
no test coverage detected