(test_file, trained_models_dir, trained_critique_dir, sequence_length,
per_gpu_eval_batch_size, language_model)
| 63 | return tokens |
| 64 | |
| 65 | def evaluate(test_file, trained_models_dir, trained_critique_dir, sequence_length, |
| 66 | per_gpu_eval_batch_size, language_model): |
| 67 | _classifier = REFINER(max_seq_length=sequence_length, |
| 68 | output_model_dir=trained_models_dir, |
| 69 | output_critique_model=trained_critique_dir, |
| 70 | cache_dir=os.path.join(DATA_FOLDER, 'pretrained'), |
| 71 | pretrained_model_name_or_path=language_model |
| 72 | ) |
| 73 | |
| 74 | print(trained_models_dir) |
| 75 | preds = _classifier.predict(test_file=test_file, |
| 76 | per_gpu_eval_batch_size=per_gpu_eval_batch_size, |
| 77 | max_generated_tokens=sequence_length) |
| 78 | |
| 79 | labels = read_labels(test_file, tag='Linear_Formula') |
| 80 | inputs = read_labels(test_file, tag='Body') |
| 81 | |
| 82 | labels = [l.lower() for l in labels] |
| 83 | preds = [p.lower() for p in preds] |
| 84 | inputs = [i for i in inputs] |
| 85 | |
| 86 | #labels = [' '.join(get_encoded_code_tokens(label)) for label in labels] |
| 87 | new_labels = [] |
| 88 | |
| 89 | with open(trained_models_dir+"/result.csv", 'w', encoding='UTF8', newline='') as outfile: |
| 90 | for index in range(len(labels)): |
| 91 | try: |
| 92 | encoded_reconstr_code = get_encoded_code_tokens(labels[index]) |
| 93 | except: |
| 94 | print("Error related to brackets", labels[index]) |
| 95 | continue |
| 96 | label = ' '.join(encoded_reconstr_code) |
| 97 | new_labels.append(labels[index]) |
| 98 | outfile.write(inputs[index] +'\t'+ preds[index] +'\t'+labels[index]+'\t'+ "yes" +'\n') |
| 99 | |
| 100 | index = 0 |
| 101 | sub_error = 0 |
| 102 | c_hyp = [tokenize_for_bleu_eval(s.lower()) for s in preds] |
| 103 | c_ref = [tokenize_for_bleu_eval(s.lower()) for s in new_labels] |
| 104 | |
| 105 | for h, r in zip(c_hyp, c_ref): |
| 106 | if h != r: |
| 107 | if 'substract' in r and 'add' not in r and 'multiply' not in r and 'divide' not in r: |
| 108 | sub_error +=1 |
| 109 | print(sub_error) |
| 110 | print(str(inputs[index]), h, r, "no", '\n') |
| 111 | |
| 112 | index += 1 |
| 113 | |
| 114 | eval_results = calculate_bleu_from_lists(gold_texts=new_labels, predicted_texts=preds) |
| 115 | print(eval_results) |
| 116 | |
| 117 | return eval_results |
| 118 | |
| 119 | def parse_args(): |
| 120 | parser = argparse.ArgumentParser(description='Critique T5') |
no test coverage detected