33# Data: 17/10/18
44# Brief: 预测
55
6- import itertools
6+
7+ import os
8+ import sys
9+
710import paddle .v2 as paddle
8- import reader
9- from network import DSSM
10- from utils import logger , ModelArch , ModelType , load_dic
11+ import numpy as np
1112import config
13+ import reader
14+ from network import dssm_lm
15+ from utils import logger , load_dict , load_reverse_dict
16+
17+
18+ def infer (model_path , dic_path , infer_path , prediction_output_path , rnn_type = "gru" , batch_size = 1 ):
19+ logger .info ("begin to predict..." )
20+ # check files
21+ assert os .path .exists (model_path ), "trained model not exits."
22+ assert os .path .exists (dic_path ), " word dictionary file not exist."
23+ assert os .path .exists (infer_path ), "infer file not exist."
24+
25+ logger .info ("load word dictionary." )
26+ word_dict = load_dict (dic_path )
27+ word_reverse_dict = load_reverse_dict (dic_path )
28+ logger .info ("dictionary size = %d" % (len (word_dict )))
1229
13- paddle .init (use_gpu = False , trainer_count = 1 )
30+ try :
31+ word_dict ["<unk>" ]
32+ except KeyError :
33+ logger .fatal ("the word dictionary must contain <unk> token." )
34+ sys .exit (- 1 )
1435
36+ # initialize PaddlePaddle
37+ paddle .init (use_gpu = config .use_gpu , trainer_count = config .num_workers )
1538
16- class Inferer (object ):
17- def __init__ (self , model_path ):
18- logger .info ("create DSSM model" )
19- self .source_dic_path = config .config ["source_dic_path" ]
20- self .target_dic_path = config .config ["target_dic_path" ]
21- dnn_dims = config .config ["dnn_dims" ]
22- layer_dims = [int (i ) for i in dnn_dims .split (',' )]
23- model_arch = ModelArch (config .config ["model_arch" ])
24- share_semantic_generator = config .config ["share_network_between_source_target" ]
25- share_embed = config .config ["share_embed" ]
26- class_num = config .config ["class_num" ]
27- prediction = DSSM (
28- dnn_dims = layer_dims ,
29- vocab_sizes = [len (load_dic (path )) for path in [self .source_dic_path , self .target_dic_path ]],
30- model_arch = model_arch ,
31- share_semantic_generator = share_semantic_generator ,
32- class_num = class_num ,
33- share_embed = share_embed ,
34- is_infer = True )()
39+ # load parameter
40+ logger .info ("load model parameters from %s " % model_path )
41+ parameters = paddle .parameters .Parameters .from_tar (
42+ open (model_path , "r" ))
3543
36- # load parameter
37- logger .info ("load model parameters from %s " % model_path )
38- self .parameters = paddle .parameters .Parameters .from_tar (
39- open (model_path , "r" ))
40- self .inferer = paddle .inference .Inference (
41- output_layer = prediction , parameters = self .parameters )
44+ # load the trained model
45+ prediction = dssm_lm (
46+ vocab_sizes = [len (word_dict ), len (word_dict )],
47+ emb_dim = config .emb_dim ,
48+ hidden_size = config .hidden_size ,
49+ stacked_rnn_num = config .stacked_rnn_num ,
50+ rnn_type = rnn_type ,
51+ share_semantic_generator = config .share_semantic_generator ,
52+ share_embed = config .share_embed ,
53+ is_infer = True )
54+ inferer = paddle .inference .Inference (
55+ output_layer = prediction , parameters = parameters )
56+ feeding = {"left_input" : 0 , "left_target" : 1 , "right_input" : 2 , "right_target" : 3 }
4257
43- def infer ( self , data_path ):
44- logger . info ( "infer data..." )
45- dataset = reader . Dataset ( train_paths = data_path ,
46- test_paths = None ,
47- source_dic_path = self . source_dic_path ,
48- target_dic_path = self . target_dic_path )
49- infer_reader = paddle . batch ( dataset . infer , batch_size = 1000 )
50- prediction_output_path = config . config [ "prediction_output_path" ]
51- logger .warning ("write prediction to %s" % prediction_output_path )
52- with open (prediction_output_path , "w" )as f :
53- for id , batch in enumerate (infer_reader ()):
54- res = self . inferer . infer ( input = batch )
55- prediction = [ " " .join (map ( str , x )) for x in res ]
56- assert len ( batch ) == len ( prediction ), ( "predict error, %d inputs,"
57- "but %d predictions" ) % (len ( batch ), len ( prediction ))
58- f .write (" \n " . join ( map ( str , prediction )) + "\n " )
58+ logger . info ( "infer data..." )
59+ # define reader
60+ reader_args = {
61+ "file_path" : infer_path ,
62+ "word_dict" : word_dict ,
63+ "is_infer" : True ,
64+ }
65+ infer_reader = paddle . batch ( reader . rnn_reader ( ** reader_args ), batch_size = batch_size )
66+ logger .warning ("output prediction to %s" % prediction_output_path )
67+ with open (prediction_output_path , "w" )as f :
68+ for id , item in enumerate (infer_reader ()):
69+ left_text = " " . join ([ word_reverse_dict [ id ] for id in item [ 0 ][ 0 ]] )
70+ right_text = " " .join ([ word_reverse_dict [ id ] for id in item [ 0 ][ 2 ]])
71+ probs = inferer . infer ( input = item , field = [ "value" ], feeding = feeding )
72+ f . write ( "%f \t %f \t %s \t %s" % (probs [ 0 ], probs [ 1 ], left_text , right_text ))
73+ f .write ("\n " )
5974
6075
61- if __name__ == '__main__' :
62- model_path = config .config ["model_path" ]
63- infer_data_paths = config .config ["infer_data_paths" ]
64- inferer = Inferer (model_path )
65- inferer .infer (infer_data_paths )
76+ if __name__ == "__main__" :
77+ infer (model_path = config .model_path ,
78+ dic_path = config .dic_path ,
79+ infer_path = config .infer_path ,
80+ prediction_output_path = config .prediction_output_path ,
81+ rnn_type = config .rnn_type )
0 commit comments