(_)
| 31 | |
| 32 | #1.load data(X:list of lint,y:int). 2.create session. 3.feed data. 4.training (5.validation) ,(6.prediction) |
| 33 | def 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 |
| 83 | def get_label_using_logits(logits,vocabulary_index2word_label,top_number=5): |
nothing calls this directly
no test coverage detected