(features, mode, params)
| 156 | |
| 157 | |
| 158 | def encoding_graph(features, mode, params): |
| 159 | if mode != "train": |
| 160 | params.residual_dropout = 0.0 |
| 161 | params.attention_dropout = 0.0 |
| 162 | params.relu_dropout = 0.0 |
| 163 | params.label_smoothing = 0.0 |
| 164 | |
| 165 | dtype = tf.get_variable_scope().dtype |
| 166 | hidden_size = params.hidden_size |
| 167 | src_seq = features["source"] |
| 168 | src_len = features["source_length"] |
| 169 | src_mask = tf.sequence_mask(src_len, |
| 170 | maxlen=tf.shape(features["source"])[1], |
| 171 | dtype=dtype or tf.float32) |
| 172 | |
| 173 | svocab = params.vocabulary["source"] |
| 174 | src_vocab_size = len(svocab) |
| 175 | initializer = tf.random_normal_initializer(0.0, params.hidden_size ** -0.5) |
| 176 | |
| 177 | if params.shared_source_target_embedding: |
| 178 | src_embedding = tf.get_variable("weights", |
| 179 | [src_vocab_size, hidden_size], |
| 180 | initializer=initializer) |
| 181 | else: |
| 182 | src_embedding = tf.get_variable("source_embedding", |
| 183 | [src_vocab_size, hidden_size], |
| 184 | initializer=initializer) |
| 185 | |
| 186 | bias = tf.get_variable("bias", [hidden_size]) |
| 187 | |
| 188 | inputs = tf.gather(src_embedding, src_seq) |
| 189 | |
| 190 | if params.multiply_embedding_mode == "sqrt_depth": |
| 191 | inputs = inputs * (hidden_size ** 0.5) |
| 192 | |
| 193 | inputs = inputs * tf.expand_dims(src_mask, -1) |
| 194 | |
| 195 | encoder_input = tf.nn.bias_add(inputs, bias) |
| 196 | enc_attn_bias = layers.attention.attention_bias(src_mask, "masking", |
| 197 | dtype=dtype) |
| 198 | if params.position_info_type == 'absolute': |
| 199 | encoder_input = layers.attention.add_timing_signal(encoder_input) |
| 200 | |
| 201 | if params.residual_dropout: |
| 202 | keep_prob = 1.0 - params.residual_dropout |
| 203 | encoder_input = tf.nn.dropout(encoder_input, keep_prob) |
| 204 | |
| 205 | encoder_output = transformer_encoder(encoder_input, enc_attn_bias, params) |
| 206 | |
| 207 | return encoder_output |
| 208 | |
| 209 | |
| 210 | def decoding_graph(features, state, mode, params): |
no test coverage detected