MCPcopy Create free account
hub / github.com/brightmart/text_classification / do_eval

Function do_eval

a02_TextCNN/p7_TextCNN_train.py:131–156  ·  view source on GitHub ↗
(sess, textCNN, evalX, evalY, num_classes)

Source from the content-addressed store, hash-verified

129
130# 在验证集上做验证,报告损失、精确度
131def do_eval(sess, textCNN, evalX, evalY, num_classes):
132 evalX = evalX[0:3000]
133 evalY = evalY[0:3000]
134 number_examples = len(evalX)
135 eval_loss, eval_counter, eval_f1_score, eval_p, eval_r = 0.0, 0, 0.0, 0.0, 0.0
136 batch_size = 1
137 predict = []
138
139 for start, end in zip(range(0, number_examples, batch_size), range(batch_size, number_examples + batch_size, batch_size)):
140 ''' evaluation in one batch '''
141 feed_dict = {textCNN.input_x: evalX[start:end],
142 textCNN.input_y_multilabel: evalY[start:end],
143 textCNN.dropout_keep_prob: 1.0,
144 textCNN.is_training_flag: False}
145 current_eval_loss, logits = sess.run(
146 [textCNN.loss_val, textCNN.logits], feed_dict)
147 predict = [*predict, np.argmax(np.array(logits[0]))]
148 eval_loss += current_eval_loss
149 eval_counter += 1
150 evalY = [np.argmax(ii) for ii in evalY]
151
152 if not FLAGS.multi_label_flag:
153 predict = [int(ii > 0.5) for ii in predict]
154 _, _, f1_macro, f1_micro, _ = fastF1(predict, evalY, num_classes)
155 f1_score = (f1_micro + f1_macro) / 2.0
156 return eval_loss / float(eval_counter), f1_score, f1_micro, f1_macro
157
158@jit
159def fastF1(result: list, predict: list, num_classes: int):

Callers 1

mainFunction · 0.70

Calls 1

fastF1Function · 0.85

Tested by

no test coverage detected