| 101 | |
| 102 | |
| 103 | def _decoder(cell, inputs, memory, sequence_length, initial_state, dtype=None, |
| 104 | scope=None): |
| 105 | # Assume that the underlying cell is GRUCell-like |
| 106 | batch = tf.shape(inputs)[0] |
| 107 | time_steps = tf.shape(inputs)[1] |
| 108 | dtype = dtype or inputs.dtype |
| 109 | output_size = cell.output_size |
| 110 | zero_output = tf.zeros([batch, output_size], dtype) |
| 111 | zero_value = tf.zeros([batch, memory.shape[-1].value], dtype) |
| 112 | |
| 113 | with tf.variable_scope(scope or "decoder", dtype=dtype): |
| 114 | inputs = tf.transpose(inputs, [1, 0, 2]) |
| 115 | mem_mask = tf.sequence_mask(sequence_length["source"], |
| 116 | maxlen=tf.shape(memory)[1], |
| 117 | dtype=dtype) |
| 118 | bias = layers.attention.attention_bias(mem_mask, "masking", |
| 119 | dtype=dtype) |
| 120 | bias = tf.squeeze(bias, axis=[1, 2]) |
| 121 | cache = layers.attention.attention(None, memory, None, output_size) |
| 122 | |
| 123 | input_ta = tf.TensorArray(dtype, time_steps, |
| 124 | tensor_array_name="input_array") |
| 125 | output_ta = tf.TensorArray(dtype, time_steps, |
| 126 | tensor_array_name="output_array") |
| 127 | value_ta = tf.TensorArray(dtype, time_steps, |
| 128 | tensor_array_name="value_array") |
| 129 | alpha_ta = tf.TensorArray(dtype, time_steps, |
| 130 | tensor_array_name="alpha_array") |
| 131 | input_ta = input_ta.unstack(inputs) |
| 132 | initial_state = layers.nn.linear(initial_state, output_size, True, |
| 133 | False, scope="s_transform") |
| 134 | initial_state = tf.tanh(initial_state) |
| 135 | |
| 136 | def loop_func(t, out_ta, att_ta, val_ta, state, cache_key): |
| 137 | inp_t = input_ta.read(t) |
| 138 | results = layers.attention.attention(state, memory, bias, |
| 139 | output_size, |
| 140 | cache={"key": cache_key}) |
| 141 | alpha = results["weight"] |
| 142 | context = results["value"] |
| 143 | cell_input = [inp_t, context] |
| 144 | cell_output, new_state = cell(cell_input, state) |
| 145 | cell_output = _copy_through(t, sequence_length["target"], |
| 146 | zero_output, cell_output) |
| 147 | new_state = _copy_through(t, sequence_length["target"], state, |
| 148 | new_state) |
| 149 | new_value = _copy_through(t, sequence_length["target"], zero_value, |
| 150 | context) |
| 151 | |
| 152 | out_ta = out_ta.write(t, cell_output) |
| 153 | att_ta = att_ta.write(t, alpha) |
| 154 | val_ta = val_ta.write(t, new_value) |
| 155 | cache_key = tf.identity(cache_key) |
| 156 | return t + 1, out_ta, att_ta, val_ta, new_state, cache_key |
| 157 | |
| 158 | time = tf.constant(0, dtype=tf.int32, name="time") |
| 159 | loop_vars = (time, output_ta, alpha_ta, value_ta, initial_state, |
| 160 | cache["key"]) |