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

Function main

a07_Transformer/a2_predict_classification.py:44–95  ·  view source on GitHub ↗
(_)

Source from the content-addressed store, hash-verified

42_PAD="_PAD"
43
44def 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
98def get_label_using_logits_batch(question_id_sublist, logits_batch, vocabulary_index2word_label, f, top_number=5):

Callers

nothing calls this directly

Calls 6

create_voabularyFunction · 0.90
create_voabulary_labelFunction · 0.90
load_final_test_dataFunction · 0.90
load_data_predictFunction · 0.90
TransformerClass · 0.90

Tested by

no test coverage detected