(features, mode, params)
| 16 | |
| 17 | |
| 18 | def model_graph(features, mode, params): |
| 19 | src_vocab_size = len(params.vocabulary["source"]) |
| 20 | tgt_vocab_size = len(params.vocabulary["target"]) |
| 21 | dtype = tf.get_variable_scope().dtype |
| 22 | |
| 23 | src_seq = features["source"] |
| 24 | tgt_seq = features["target"] |
| 25 | |
| 26 | if params.reverse_source: |
| 27 | src_seq = tf.reverse_sequence(src_seq, seq_dim=1, |
| 28 | seq_lengths=features["source_length"]) |
| 29 | |
| 30 | with tf.device("/cpu:0"): |
| 31 | with tf.variable_scope("source_embedding"): |
| 32 | src_emb = tf.get_variable("embedding", |
| 33 | [src_vocab_size, params.embedding_size]) |
| 34 | src_bias = tf.get_variable("bias", [params.embedding_size]) |
| 35 | src_inputs = tf.nn.embedding_lookup(src_emb, src_seq) |
| 36 | |
| 37 | with tf.variable_scope("target_embedding"): |
| 38 | tgt_emb = tf.get_variable("embedding", |
| 39 | [tgt_vocab_size, params.embedding_size]) |
| 40 | tgt_bias = tf.get_variable("bias", [params.embedding_size]) |
| 41 | tgt_inputs = tf.nn.embedding_lookup(tgt_emb, tgt_seq) |
| 42 | |
| 43 | src_inputs = tf.nn.bias_add(src_inputs, src_bias) |
| 44 | tgt_inputs = tf.nn.bias_add(tgt_inputs, tgt_bias) |
| 45 | |
| 46 | if params.dropout and not params.use_variational_dropout: |
| 47 | src_inputs = tf.nn.dropout(src_inputs, 1.0 - params.dropout) |
| 48 | tgt_inputs = tf.nn.dropout(tgt_inputs, 1.0 - params.dropout) |
| 49 | |
| 50 | cell_enc = [] |
| 51 | cell_dec = [] |
| 52 | |
| 53 | for _ in range(params.num_hidden_layers): |
| 54 | if params.rnn_cell == "LSTMCell": |
| 55 | cell_e = tf.nn.rnn_cell.BasicLSTMCell(params.hidden_size) |
| 56 | cell_d = tf.nn.rnn_cell.BasicLSTMCell(params.hidden_size) |
| 57 | elif params.rnn_cell == "GRUCell": |
| 58 | cell_e = tf.nn.rnn_cell.GRUCell(params.hidden_size) |
| 59 | cell_d = tf.nn.rnn_cell.GRUCell(params.hidden_size) |
| 60 | else: |
| 61 | raise ValueError("%s not supported" % params.rnn_cell) |
| 62 | |
| 63 | cell_e = tf.nn.rnn_cell.DropoutWrapper( |
| 64 | cell_e, |
| 65 | output_keep_prob=1.0 - params.dropout, |
| 66 | variational_recurrent=params.use_variational_dropout, |
| 67 | input_size=params.embedding_size, |
| 68 | dtype=dtype |
| 69 | ) |
| 70 | cell_d = tf.nn.rnn_cell.DropoutWrapper( |
| 71 | cell_d, |
| 72 | output_keep_prob=1.0 - params.dropout, |
| 73 | variational_recurrent=params.use_variational_dropout, |
| 74 | input_size=params.embedding_size, |
| 75 | dtype=dtype |
no outgoing calls
no test coverage detected