| 177 | |
| 178 | # 在验证集上做验证,报告损失、精确度 |
| 179 | def do_eval(sess,model,evalX,evalY,batch_size,vocabulary_index2word_label,eval_decoder_input=None): |
| 180 | #ii=0 |
| 181 | number_examples=len(evalX) |
| 182 | eval_loss,eval_acc,eval_counter=0.0,0.0,0 |
| 183 | for start,end in zip(range(0,number_examples,batch_size),range(batch_size,number_examples,batch_size)): |
| 184 | feed_dict = {model.query: evalX[start:end],model.story:np.expand_dims(evalX[start:end],axis=1), model.dropout_keep_prob: 1} |
| 185 | if not FLAGS.multi_label_flag: |
| 186 | feed_dict[model.answer_single] = evalY[start:end] |
| 187 | else: |
| 188 | feed_dict[model.answer_multilabel] = evalY[start:end] |
| 189 | curr_eval_loss, logits,curr_eval_acc,pred= sess.run([model.loss_val,model.logits,model.accuracy,model.predictions],feed_dict)#curr_eval_acc--->textCNN.accuracy |
| 190 | eval_loss,eval_acc,eval_counter=eval_loss+curr_eval_loss,eval_acc+curr_eval_acc,eval_counter+1 |
| 191 | #if ii<20: |
| 192 | #print("1.evalX[start:end]:",evalX[start:end]) |
| 193 | #print("2.evalY[start:end]:", evalY[start:end]) |
| 194 | #print("3.pred:",pred) |
| 195 | #ii=ii+1 |
| 196 | return eval_loss/float(eval_counter),eval_acc/float(eval_counter) |
| 197 | |
| 198 | #从logits中取出前五 get label using logits |
| 199 | def get_label_using_logits(logits,vocabulary_index2word_label,top_number=1): |