MCPcopy Create free account
hub / github.com/debjitpaul/refiner / evaluate

Function evaluate

src/scripts/train_refiner.py:65–117  ·  view source on GitHub ↗
(test_file, trained_models_dir, trained_critique_dir, sequence_length,
             per_gpu_eval_batch_size, language_model)

Source from the content-addressed store, hash-verified

63 return tokens
64
65def 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
119def parse_args():
120 parser = argparse.ArgumentParser(description='Critique T5')

Callers 1

mainFunction · 0.70

Calls 6

predictMethod · 0.95
REFINERClass · 0.90
read_labelsFunction · 0.90
get_encoded_code_tokensFunction · 0.90
tokenize_for_bleu_evalFunction · 0.70

Tested by

no test coverage detected