tf_bi_lstm.py 1.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849
  1. import tensorflow as tf
  2. from tensorflow.contrib.rnn import LSTMCell
  3. from tensorflow.contrib.rnn import MultiRNNCell
  4. class LstmBase:
  5. """
  6. build rnn cell
  7. """
  8. def build_rnn(self, hidden_size, num_layes):
  9. cells = []
  10. for i in range(num_layes):
  11. cell = LSTMCell(num_units=hidden_size,
  12. state_is_tuple=True,
  13. initializer=tf.random_uniform_initializer(-0.25, 0.25))
  14. cells.append(cell)
  15. cells = MultiRNNCell(cells, state_is_tuple=True)
  16. return cells
  17. class BiLstm(LstmBase):
  18. """
  19. define the lstm
  20. """
  21. def __init__(self, scope_name, hidden_size, num_layers):
  22. super(BiLstm, self).__init__()
  23. assert hidden_size % 2 == 0
  24. hidden_size /= 2
  25. self.fw_rnns = []
  26. self.bw_rnns = []
  27. for i in range(num_layers):
  28. self.fw_rnns.append(self.build_rnn(hidden_size, 1))
  29. self.bw_rnns.append(self.build_rnn(hidden_size, 1))
  30. self.scope_name = scope_name
  31. def __call__(self, input, input_len):
  32. for idx, (fw_rnn, bw_rnn) in enumerate(zip(self.fw_rnns, self.bw_rnns)):
  33. scope_name = '{}_{}'.format(self.scope_name, idx)
  34. ctx, _ = tf.nn.bidirectional_dynamic_rnn(
  35. fw_rnn, bw_rnn, input, sequence_length=input_len,
  36. dtype=tf.float32, time_major=False,
  37. scope=scope_name
  38. )
  39. input = tf.concat(ctx, -1)
  40. ctx = input
  41. return ctx