MCPcopy Create free account
hub / github.com/XL2248/MSCTD / _decoder

Function _decoder

src_code/thumt1_code/thumt/models/rnnsearch.py:103–184  ·  view source on GitHub ↗
(cell, inputs, memory, sequence_length, initial_state, dtype=None,
             scope=None)

Source from the content-addressed store, hash-verified

101
102
103def _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"])

Callers 1

model_graphFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected