构造模型
(feed_previous=False)
| 114 | |
| 115 | |
| 116 | def 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 | |
| 162 | def train(): |
no test coverage detected