Skip to content

Commit 983f293

Browse files
author
xuming06
committed
update dssm. xuming 20171020
1 parent 4707eb0 commit 983f293

4 files changed

Lines changed: 82 additions & 17 deletions

File tree

16paddle/dssm/config.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,8 +19,12 @@
1919
"num_workers": 4, # num worker threads, default 1
2020
"use_gpu": False, # use GPU devices
2121
"class_num": 2, # number of categories for classification task
22-
"model_output_prefix": "./output/", # prefix of the path for model to store
22+
"model_output_prefix": "./data/output/", # prefix of the path for model to store
2323
"num_batches_to_log": 100, # log
2424
"num_batches_to_test": 200, # batches to test
2525
"num_batches_to_save_model": 400, # number of batches to output model, (default: 400)
26+
27+
"prediction_output_path": "./data/output/prediction.txt", # output prediction file
28+
"model_path": "./data/output/dssm_classification_rnn_pass_00009.tar", # saved model path
29+
"infer_data_paths": ["./data/classification/test/right.txt"], # infer data path
2630
}
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
汽车 驾驶 驾校 培训:
1+
汽车 驾驶 驾校 培训

16paddle/dssm/infer.py

Lines changed: 62 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,65 @@
11
# -*- coding: utf-8 -*-
22
# Author: XuMing <[email protected]>
33
# Data: 17/10/18
4-
# Brief:
4+
# Brief: 预测
5+
6+
import itertools
7+
import paddle.v2 as paddle
8+
import reader
9+
from network import DSSM
10+
from utils import logger, ModelArch, ModelType, load_dic
11+
import config
12+
13+
paddle.init(use_gpu=False, trainer_count=1)
14+
15+
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)()
35+
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)
42+
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")
59+
60+
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)

16paddle/dssm/train.py

Lines changed: 14 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -5,25 +5,24 @@
55

66

77
import 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+
119
import config
10+
import reader
11+
from network import DSSM
12+
from utils import load_dic, logger, ModelType, ModelArch, display_args
1213

1314

1415
def 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

Comments
 (0)