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