构造模型
(feed_previous=False)
| 54 | |
| 55 | |
| 56 | def 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 | |
| 97 | def train(): |
no test coverage detected