55
66
77import paddle .v2 as paddle
8- from network import DSSM
9- import reader
10- from utils import TaskType , load_dic , logger , ModelType , ModelArch , display_args
8+
119import config
10+ import reader
11+ from network import DSSM
12+ from utils import load_dic , logger , ModelType , ModelArch , display_args
1213
1314
1415def train (train_data_paths = None ,
1516 test_data_paths = None ,
1617 source_dic_path = None ,
1718 target_dic_path = None ,
18- model_type = ModelType .create_classification (),
1919 model_arch = ModelArch .create_rnn (),
2020 batch_size = 10 ,
2121 num_passes = 10 ,
2222 share_semantic_generator = False ,
2323 share_embed = False ,
2424 class_num = 2 ,
2525 num_workers = 1 ,
26- dnn_dims = "256,128,64,32" ,
2726 use_gpu = False ):
2827 """
2928 train DSSM
@@ -33,7 +32,7 @@ def train(train_data_paths=None,
3332 default_test_paths = ["./data/classification/test/right.txt" ,
3433 "./data/classification/test/wrong.txt" ]
3534 default_dic_path = "./data/vocab.txt"
36- layer_dims = [int (i ) for i in dnn_dims .split (',' )]
35+ layer_dims = [int (i ) for i in config . config [ ' dnn_dims' ] .split (',' )]
3736 use_default_data = not train_data_paths
3837 if use_default_data :
3938 train_data_paths = default_train_paths
@@ -58,7 +57,6 @@ def train(train_data_paths=None,
5857 cost , prediction , label = DSSM (
5958 dnn_dims = layer_dims ,
6059 vocab_sizes = [len (load_dic (path )) for path in [source_dic_path , target_dic_path ]],
61- model_type = model_type ,
6260 model_arch = model_arch ,
6361 share_semantic_generator = share_semantic_generator ,
6462 class_num = class_num ,
@@ -93,22 +91,25 @@ def _event_handler(event):
9391 # test model
9492 if event .batch_id > 0 and event .batch_id % config .config ['num_batches_to_test' ] == 0 :
9593 if test_reader is not None :
96- if model_type .is_classification ():
97- result = trainer .test (reader = test_reader , feeding = feeding )
98- logger .info ("Test at Pass %d, %s" % (event .pass_id , result .metrics ))
99- else :
100- result = None
94+ result = trainer .test (reader = test_reader , feeding = feeding )
95+ logger .info ("Test at Pass %d, %s" % (event .pass_id , result .metrics ))
10196
10297 # save model
10398 if event .batch_id > 0 and event .batch_id % config .config ['num_batches_to_save_model' ] == 0 :
104- model_desc = "classification_{arch}" .format (arch = str (config . config [ ' model_arch' ] ))
99+ model_desc = "classification_{arch}" .format (arch = str (model_arch ))
105100 with open ("%sdssm_%s_pass_%05d.tar" %
106101 (config .config ['model_output_prefix' ], model_desc ,
107102 event .pass_id ), "w" ) as f :
108103 parameters .to_tar (f )
109104 logger .info ("save model: %sdssm_%s_pass_%05d.tar" %
110105 (config .config ['model_output_prefix' ], model_desc , event .pass_id ))
111106
107+ # if isinstance(event, paddle.event.EndPass):
108+ # result = trainer.test(reader=test_reader, feeding=feeding)
109+ # logger.info("Test with pass %d, %s" % (event.pass_id, result.metrics))
110+ # with open("./data/output/endpass/dssm_params_pass" + str(event.pass_id) + ".tar", "w") as f:
111+ # parameters.to_tar(f)
112+
112113 trainer .train (reader = train_reader ,
113114 event_handler = _event_handler ,
114115 feeding = feeding ,
@@ -127,7 +128,6 @@ def _event_handler(event):
127128 num_passes = config .config ["num_passes" ],
128129 share_semantic_generator = config .config ["share_network_between_source_target" ],
129130 share_embed = config .config ["share_embed" ],
130- dnn_dims = config .config ["dnn_dims" ],
131131 class_num = config .config ["class_num" ],
132132 num_workers = config .config ["num_workers" ],
133133 use_gpu = config .config ["use_gpu" ])
0 commit comments