do online prediction. each time make prediction for one instance. you can change to a batch if you want. :param line: a list. element is: [dummy_label,text_a,text_b] :return:
(line)
| 352 | saver.restore(sess, tf.train.latest_checkpoint(FLAGS.init_checkpoint)) |
| 353 | |
| 354 | def predict_online(line): |
| 355 | """ |
| 356 | do online prediction. each time make prediction for one instance. |
| 357 | you can change to a batch if you want. |
| 358 | :param line: a list. element is: [dummy_label,text_a,text_b] |
| 359 | :return: |
| 360 | """ |
| 361 | label = line[0] #tokenization.convert_to_unicode(line[0]) # this should compatible with format you defined in processor. |
| 362 | text_a = line[1] #tokenization.convert_to_unicode(line[1]) |
| 363 | text_b = line[2] #tokenization.convert_to_unicode(line[2]) |
| 364 | example= InputExample(guid=0, text_a=text_a, text_b=text_b, label=label) |
| 365 | feature = convert_single_example(0, example, label_list,FLAGS.max_seq_length, tokenizer) |
| 366 | input_ids = np.reshape([feature.input_ids],(1,FLAGS.max_seq_length)) |
| 367 | input_mask = np.reshape([feature.input_mask],(1,FLAGS.max_seq_length)) |
| 368 | segment_ids = np.reshape([feature.segment_ids],(FLAGS.max_seq_length)) |
| 369 | label_ids =[feature.label_id] |
| 370 | |
| 371 | global graph |
| 372 | with graph.as_default(): |
| 373 | feed_dict = {input_ids_p: input_ids, input_mask_p: input_mask,segment_ids_p:segment_ids,label_ids_p:label_ids} |
| 374 | possibility = sess.run([probabilities], feed_dict) |
| 375 | possibility=possibility[0][0] # get first label |
| 376 | label_index=np.argmax(possibility) |
| 377 | label_predict=index2label[label_index] |
| 378 | #print("label_predict:",label_predict,";possibility:",possibility) |
| 379 | return label_predict,possibility |
| 380 | |
| 381 | if __name__ == "__main__": |
| 382 | example=['0','\u5165\u804c\u4e00\u5e74\u534a\u672a\u7b7e\u52b3\u52a8\u5408\u540c\u5c0f\u83f2\u6bd5\u4e1a\u4e8e\u67d0\u62a4\u6821\uff0c\u548c\u5176\u4ed6\u7684\u9ad8\u6821\u6bd5\u4e1a\u751f\u4e00\u6837\uff0c\u5979\u4e5f\u5f00\u59cb\u7740\u624b\u627e\u5de5\u4f5c\u3002\u5f88\u5feb\uff0c\u4e00\u5bb6\u6c11\u529e\u533b\u9662\u901a\u8fc7\u67d0\u62db\u8058\u7f51\u7ad9\u627e\u5230\u5c0f\u83f2\uff0c\u901a\u8fc7\u9762\u8bd5\u540e\uff0c\u5c0f\u83f2\u4fbf\u5f00\u59cb\u4e86\u81ea\u5df1\u7684\u804c\u573a\u751f\u6daf\u3002\u8f6c\u773c\u6bd5\u4e1a\u5de5\u4f5c\u8fd1\u4e00\u5e74\uff0c\u533b\u9662\u4ecd\u8fdf\u8fdf\u4e0d\u4e0e\u5176\u7b7e\u8ba2\u52b3\u52a8\u5408\u540c\uff0c\u5c0f\u83f2\u4e0e\u5355\u4f4d\u591a\u6b21\u6c9f\u901a\u534f\u5546\u672a\u679c\uff0c\u65e0\u5948\u5c06\u533b\u9662\u8bc9\u81f3\u6cd5\u9662\u000d\u000a','\u652f\u4ed8\u5de5\u8d44'] |
no test coverage detected