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

Function main

a01_FastText/p5_fastTextB_predict_multilabel.py:33–80  ·  view source on GitHub ↗
(_)

Source from the content-addressed store, hash-verified

31
32#1.load data(X:list of lint,y:int). 2.create session. 3.feed data. 4.training (5.validation) ,(6.prediction)
33def main(_):
34 # 1.load data with vocabulary of words and labels
35 vocabulary_word2index, vocabulary_index2word = create_voabulary()
36 vocab_size = len(vocabulary_word2index)
37 print("vocab_size:",vocab_size)
38 #iii=0
39 #iii/0
40 vocabulary_word2index_label,vocabulary_index2word_label = create_voabulary_label()
41 questionid_question_lists=load_final_test_data(FLAGS.predict_source_file) #TODO
42 test= load_data_predict(vocabulary_word2index,vocabulary_word2index_label,questionid_question_lists) #TODO
43 testX=[]
44 question_id_list=[]
45 for tuple in test:
46 question_id,question_string_list=tuple
47 question_id_list.append(question_id)
48 testX.append(question_string_list)
49
50 # 2.Data preprocessing: Sequence padding
51 print("start padding....")
52 testX2 = pad_sequences(testX, maxlen=FLAGS.sentence_len, value=0.) # padding to max length
53 print("end padding...")
54
55 # 3.create session.
56 config=tf.ConfigProto()
57 config.gpu_options.allow_growth=True
58 with tf.Session(config=config) as sess:
59 # 4.Instantiate Model
60 fast_text=fastText(FLAGS.label_size, FLAGS.learning_rate, FLAGS.batch_size, FLAGS.decay_steps, FLAGS.decay_rate,FLAGS.num_sampled,FLAGS.sentence_len,vocab_size,FLAGS.embed_size,FLAGS.is_training)
61 saver=tf.train.Saver()
62 if os.path.exists(FLAGS.ckpt_dir+"checkpoint"):
63 print("Restoring Variables from Checkpoint")
64 saver.restore(sess,tf.train.latest_checkpoint(FLAGS.ckpt_dir))
65 else:
66 print("Can't find the checkpoint.going to stop")
67 return
68 # 5.feed data, to get logits
69 number_of_training_data=len(testX2);print("number_of_training_data:",number_of_training_data)
70 batch_size=1
71 index=0
72 predict_target_file_f = codecs.open(FLAGS.predict_target_file, 'a', 'utf8')
73 for start, end in zip(range(0, number_of_training_data, batch_size),range(batch_size, number_of_training_data+1, batch_size)):
74 logits=sess.run(fast_text.logits,feed_dict={fast_text.sentence:testX2[start:end]}) #'shape of logits:', ( 1, 1999)
75 # 6. get lable using logtis
76 predicted_labels=get_label_using_logits(logits[0],vocabulary_index2word_label)
77 # 7. write question id and labels to file system.
78 write_question_id_with_labels(question_id_list[index],predicted_labels,predict_target_file_f)
79 index=index+1
80 predict_target_file_f.close()
81
82# get label using logits
83def get_label_using_logits(logits,vocabulary_index2word_label,top_number=5):

Callers

nothing calls this directly

Calls 6

get_label_using_logitsFunction · 0.70
create_voabularyFunction · 0.50
create_voabulary_labelFunction · 0.50
load_final_test_dataFunction · 0.50
load_data_predictFunction · 0.50

Tested by

no test coverage detected