| 129 | |
| 130 | # 在验证集上做验证,报告损失、精确度 |
| 131 | def 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 |
| 159 | def fastF1(result: list, predict: list, num_classes: int): |