| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849 |
- import tensorflow as tf
- from tensorflow.contrib.rnn import LSTMCell
- from tensorflow.contrib.rnn import MultiRNNCell
- class LstmBase:
- """
- build rnn cell
- """
- def build_rnn(self, hidden_size, num_layes):
- cells = []
- for i in range(num_layes):
- cell = LSTMCell(num_units=hidden_size,
- state_is_tuple=True,
- initializer=tf.random_uniform_initializer(-0.25, 0.25))
- cells.append(cell)
- cells = MultiRNNCell(cells, state_is_tuple=True)
- return cells
- class BiLstm(LstmBase):
- """
- define the lstm
- """
- def __init__(self, scope_name, hidden_size, num_layers):
- super(BiLstm, self).__init__()
- assert hidden_size % 2 == 0
- hidden_size /= 2
- self.fw_rnns = []
- self.bw_rnns = []
- for i in range(num_layers):
- self.fw_rnns.append(self.build_rnn(hidden_size, 1))
- self.bw_rnns.append(self.build_rnn(hidden_size, 1))
- self.scope_name = scope_name
- def __call__(self, input, input_len):
- for idx, (fw_rnn, bw_rnn) in enumerate(zip(self.fw_rnns, self.bw_rnns)):
- scope_name = '{}_{}'.format(self.scope_name, idx)
- ctx, _ = tf.nn.bidirectional_dynamic_rnn(
- fw_rnn, bw_rnn, input, sequence_length=input_len,
- dtype=tf.float32, time_major=False,
- scope=scope_name
- )
- input = tf.concat(ctx, -1)
- ctx = input
- return ctx
|