MCPcopy Create free account
hub / github.com/chinawithfrank/ChatBotCourse / get_model

Function get_model

chatbotv4/demo.py:56–94  ·  view source on GitHub ↗

构造模型

(feed_previous=False)

Source from the content-addressed store, hash-verified

54
55
56def get_model(feed_previous=False):
57 """构造模型
58 """
59 encoder_inputs = []
60 decoder_inputs = []
61 target_weights = []
62 for i in xrange(input_seq_len):
63 encoder_inputs.append(tf.placeholder(tf.int32, shape=[None], name="encoder{0}".format(i)))
64 for i in xrange(output_seq_len + 1):
65 decoder_inputs.append(tf.placeholder(tf.int32, shape=[None], name="decoder{0}".format(i)))
66 for i in xrange(output_seq_len):
67 target_weights.append(tf.placeholder(tf.float32, shape=[None], name="weight{0}".format(i)))
68
69 # decoder_inputs左移一个时序作为targets
70 targets = [decoder_inputs[i + 1] for i in xrange(output_seq_len)]
71
72 cell = tf.contrib.rnn.BasicLSTMCell(size)
73
74 # 这里输出的状态我们不需要
75 outputs, _ = seq2seq.embedding_attention_seq2seq(
76 encoder_inputs,
77 decoder_inputs[:output_seq_len],
78 cell,
79 num_encoder_symbols=num_encoder_symbols,
80 num_decoder_symbols=num_decoder_symbols,
81 embedding_size=size,
82 output_projection=None,
83 feed_previous=feed_previous,
84 dtype=tf.float32)
85
86 # 计算加权交叉熵损失
87 loss = seq2seq.sequence_loss(outputs, targets, target_weights)
88 # 梯度下降优化器
89 opt = tf.train.GradientDescentOptimizer(learning_rate)
90 # 优化目标:让loss最小化
91 update = opt.apply_gradients(opt.compute_gradients(loss))
92 # 模型持久化
93 saver = tf.train.Saver(tf.global_variables())
94 return encoder_inputs, decoder_inputs, target_weights, outputs, loss, update, saver
95
96
97def train():

Callers 2

trainFunction · 0.70
predictFunction · 0.70

Calls 2

sequence_lossMethod · 0.80
appendMethod · 0.45

Tested by

no test coverage detected