(_)
| 42 | _PAD="_PAD" |
| 43 | |
| 44 | def main(_): |
| 45 | # 1.load data with vocabulary of words and labels |
| 46 | vocabulary_word2index, vocabulary_index2word = create_voabulary(word2vec_model_path=FLAGS.word2vec_model_path,name_scope="transformer_classification") # simple='simple' |
| 47 | vocab_size = len(vocabulary_word2index) |
| 48 | print("transformer_classification.vocab_size:", vocab_size) |
| 49 | vocabulary_word2index_label, vocabulary_index2word_label = create_voabulary_label(name_scope="transformer_classification") |
| 50 | questionid_question_lists=load_final_test_data(FLAGS.predict_source_file) |
| 51 | print("list of total questions:",len(questionid_question_lists)) |
| 52 | test= load_data_predict(vocabulary_word2index,vocabulary_word2index_label,questionid_question_lists) |
| 53 | print("list of total questions2:",len(test)) |
| 54 | testX=[] |
| 55 | question_id_list=[] |
| 56 | for tuple in test: |
| 57 | question_id,question_string_list=tuple |
| 58 | question_id_list.append(question_id) |
| 59 | testX.append(question_string_list) |
| 60 | # 2.Data preprocessing: Sequence padding |
| 61 | print("start padding....") |
| 62 | testX2 = pad_sequences(testX, maxlen=FLAGS.sequence_length, value=0.) # padding to max length |
| 63 | print("list of total questions3:", len(testX2)) |
| 64 | print("end padding...") |
| 65 | # 3.create session. |
| 66 | config=tf.ConfigProto() |
| 67 | config.gpu_options.allow_growth=True |
| 68 | with tf.Session(config=config) as sess: |
| 69 | # 4.Instantiate Model |
| 70 | model=Transformer(FLAGS.num_classes, FLAGS.learning_rate, FLAGS.batch_size, FLAGS.decay_steps, FLAGS.decay_rate, FLAGS.sequence_length, |
| 71 | vocab_size, FLAGS.embed_size,FLAGS.d_model,FLAGS.d_k,FLAGS.d_v,FLAGS.h,FLAGS.num_layer,FLAGS.is_training,l2_lambda=FLAGS.l2_lambda) |
| 72 | saver=tf.train.Saver() |
| 73 | if os.path.exists(FLAGS.ckpt_dir+"checkpoint"): |
| 74 | print("Restoring Variables from Checkpoint") |
| 75 | saver.restore(sess,tf.train.latest_checkpoint(FLAGS.ckpt_dir)) |
| 76 | else: |
| 77 | print("Can't find the checkpoint.going to stop") |
| 78 | return |
| 79 | # 5.feed data, to get logits |
| 80 | number_of_training_data=len(testX2);print("number_of_training_data:",number_of_training_data) |
| 81 | index=0 |
| 82 | predict_target_file_f = codecs.open(FLAGS.predict_target_file, 'a', 'utf8') |
| 83 | for start, end in zip(range(0, number_of_training_data, FLAGS.batch_size),range(FLAGS.batch_size, number_of_training_data+1, FLAGS.batch_size)): |
| 84 | logits=sess.run(model.logits,feed_dict={model.input_x:testX2[start:end],model.dropout_keep_prob:1}) #logits:[batch_size,self.num_classes] |
| 85 | |
| 86 | question_id_sublist=question_id_list[start:end] |
| 87 | get_label_using_logits_batch(question_id_sublist, logits, vocabulary_index2word_label, predict_target_file_f) |
| 88 | |
| 89 | # 6. get lable using logtis |
| 90 | #predicted_labels=get_label_using_logits(logits[0],vocabulary_index2word_label) |
| 91 | #print(index," ;predicted_labels:",predicted_labels) |
| 92 | # 7. write question id and labels to file system. |
| 93 | #write_question_id_with_labels(question_id_list[index],predicted_labels,predict_target_file_f) |
| 94 | index=index+1 |
| 95 | predict_target_file_f.close() |
| 96 | |
| 97 | # get label using logits |
| 98 | def get_label_using_logits_batch(question_id_sublist, logits_batch, vocabulary_index2word_label, f, top_number=5): |
nothing calls this directly
no test coverage detected