(test_file, trained_models_dir, trained_critique_dir, sequence_length,
per_gpu_eval_batch_size, language_model)
| 40 | |
| 41 | |
| 42 | def evaluate(test_file, trained_models_dir, trained_critique_dir, sequence_length, |
| 43 | per_gpu_eval_batch_size, language_model): |
| 44 | _classifier = T5LMClassifier(max_seq_length=sequence_length, |
| 45 | output_model_dir=trained_models_dir, |
| 46 | output_critique_model=trained_critique_dir, |
| 47 | cache_dir=os.path.join(DATA_FOLDER, 'pretrained'), |
| 48 | pretrained_model_name_or_path=language_model |
| 49 | ) |
| 50 | |
| 51 | print(trained_models_dir) |
| 52 | preds = _classifier.predict(test_file=test_file, |
| 53 | per_gpu_eval_batch_size=per_gpu_eval_batch_size, |
| 54 | max_generated_tokens=sequence_length) |
| 55 | |
| 56 | labels = read_labels(test_file, tag='Linear_Formula') |
| 57 | inputs = read_labels(test_file, tag='Body') |
| 58 | |
| 59 | labels = [l.lower() for l in labels] |
| 60 | preds = [p.lower() for p in preds] |
| 61 | inputs = [i for i in inputs] |
| 62 | |
| 63 | #labels = [' '.join(get_encoded_code_tokens(label)) for label in labels] |
| 64 | new_labels = [] |
| 65 | |
| 66 | with open(trained_models_dir+"/result.csv", 'w', encoding='UTF8', newline='') as outfile: |
| 67 | for index in range(len(labels)): |
| 68 | try: |
| 69 | encoded_reconstr_code = get_encoded_code_tokens(labels[index]) |
| 70 | except: |
| 71 | print("Error related to brackets", labels[index]) |
| 72 | continue |
| 73 | label = ' '.join(encoded_reconstr_code) |
| 74 | new_labels.append(labels[index]) |
| 75 | print(preds[index].strip() == labels[index].strip()) |
| 76 | if preds[index].strip() == labels[index].strip(): |
| 77 | outfile.write(inputs[index] +'\t'+ preds[index] +'\t'+labels[index]+'\t'+ "yes" +'\n') |
| 78 | else: |
| 79 | outfile.write(inputs[index] +'\t'+ preds[index] +'\t'+labels[index]+'\t'+ "no" +'\n') |
| 80 | |
| 81 | |
| 82 | eval_results = calculate_bleu_from_lists(gold_texts=new_labels, predicted_texts=preds) |
| 83 | print(eval_results) |
| 84 | |
| 85 | return eval_results |
| 86 | |
| 87 | def parse_args(): |
| 88 | parser = argparse.ArgumentParser(description='Critique T5') |
no test coverage detected