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

Function get_model

chatbotv5/demo.py:116–159  ·  view source on GitHub ↗

构造模型

(feed_previous=False)

Source from the content-addressed store, hash-verified

114
115
116def get_model(feed_previous=False):
117 """构造模型
118 """
119
120 learning_rate = tf.Variable(float(init_learning_rate), trainable=False, dtype=tf.float32)
121 learning_rate_decay_op = learning_rate.assign(learning_rate * 0.9)
122
123 encoder_inputs = []
124 decoder_inputs = []
125 target_weights = []
126 for i in xrange(input_seq_len):
127 encoder_inputs.append(tf.placeholder(tf.int32, shape=[None], name="encoder{0}".format(i)))
128 for i in xrange(output_seq_len + 1):
129 decoder_inputs.append(tf.placeholder(tf.int32, shape=[None], name="decoder{0}".format(i)))
130 for i in xrange(output_seq_len):
131 target_weights.append(tf.placeholder(tf.float32, shape=[None], name="weight{0}".format(i)))
132
133 # decoder_inputs左移一个时序作为targets
134 targets = [decoder_inputs[i + 1] for i in xrange(output_seq_len)]
135
136 cell = tf.contrib.rnn.BasicLSTMCell(size)
137
138 # 这里输出的状态我们不需要
139 outputs, _ = seq2seq.embedding_attention_seq2seq(
140 encoder_inputs,
141 decoder_inputs[:output_seq_len],
142 cell,
143 num_encoder_symbols=num_encoder_symbols,
144 num_decoder_symbols=num_decoder_symbols,
145 embedding_size=size,
146 output_projection=None,
147 feed_previous=feed_previous,
148 dtype=tf.float32)
149
150 # 计算加权交叉熵损失
151 loss = seq2seq.sequence_loss(outputs, targets, target_weights)
152 # 梯度下降优化器
153 opt = tf.train.GradientDescentOptimizer(learning_rate)
154 # 优化目标:让loss最小化
155 update = opt.apply_gradients(opt.compute_gradients(loss))
156 # 模型持久化
157 saver = tf.train.Saver(tf.global_variables())
158
159 return encoder_inputs, decoder_inputs, target_weights, outputs, loss, update, saver, learning_rate_decay_op, learning_rate
160
161
162def train():

Callers 2

trainFunction · 0.70
predictFunction · 0.70

Calls 2

sequence_lossMethod · 0.80
appendMethod · 0.45

Tested by

no test coverage detected