From 44a42a6a9835f2b84adb140cba6b7dc887c719d1 Mon Sep 17 00:00:00 2001 From: Hariharan Seshadri Date: Fri, 16 Aug 2019 10:12:46 -0700 Subject: [PATCH] Fix parsing initial hidden state in RNN (#1626) * Fix the way initial hidden state is used for reverse direction in RNN * Add test case * Updates --- onnxruntime/core/providers/cpu/rnn/rnn.cc | 8 ++- .../test/providers/cpu/rnn/rnn_op_test.cc | 54 ++++++++++++++++++- 2 files changed, 59 insertions(+), 3 deletions(-) diff --git a/onnxruntime/core/providers/cpu/rnn/rnn.cc b/onnxruntime/core/providers/cpu/rnn/rnn.cc index 4030d65a94..5c1b234d1e 100644 --- a/onnxruntime/core/providers/cpu/rnn/rnn.cc +++ b/onnxruntime/core/providers/cpu/rnn/rnn.cc @@ -181,8 +181,12 @@ Status RNN::Compute(OpKernelContext* ctx) const { const float* h_prev = nullptr; if (t == 0) { - if (initial_h != nullptr) - h_prev = initial_h->template Data(); + if (initial_h != nullptr) { + // the shape of initial_h is [num_directions, batch_size, hidden_size] + // so pick the offset (multiple of Y_frame_size == batch_size * hidden_size_) + // based on the direction + h_prev = initial_h->template Data() + (direction * Y_frame_size); + } } else { if (isReverse) h_prev = Y_buffer_data_current_frame + num_directions * Y_frame_size; diff --git a/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc b/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc index a7412d32e2..2b9e81c149 100644 --- a/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc +++ b/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc @@ -345,7 +345,7 @@ TEST(RNNTest, RNN_forward_direction_zigged_batch) { test.Run(); } -TEST(RNNTest, RNN_bidirectional) { +TEST(RNNTest, RNN_bidirectional_0) { OpTester test("RNN"); int64_t num_directions = 2, input_size = 2, hidden_size = 3, batch_size = 1, seq_length = 5; @@ -407,6 +407,58 @@ TEST(RNNTest, RNN_bidirectional) { test.Run(); } +TEST(RNNTest, RNN_bidirectional_1) { + OpTester test("RNN"); + int64_t num_directions = 2, input_size = 2, hidden_size = 2, batch_size = 1, seq_length = 1; + + test.AddAttribute("activations", vector(num_directions, "Tanh")); + test.AddAttribute("direction", "bidirectional"); + test.AddAttribute("hidden_size", hidden_size); + + std::vector X_dims = {seq_length, batch_size, input_size}; + std::vector X_data({1.0F, 1.0F}); + + test.AddInput("X", X_dims, X_data); + + std::vector W_dims = {num_directions, hidden_size, input_size}; + std::vector W_data({1.0F, 1.0F, 1.0F, 1.0F, + 1.0F, 1.0F, 1.0F, 1.0F}); + + test.AddInput("W", W_dims, W_data); + + std::vector R_dims = {num_directions, hidden_size, hidden_size}; + std::vector R_data({// forward + 1.0F, 1.0F, + 1.0F, 1.0F, + // reverse + 1.0F, 1.0F, + 1.0F, 1.0F}); + test.AddInput("R", R_dims, R_data); + + std::vector B_dims = {num_directions, 2 * hidden_size}; + std::vector B_data({0.0F, 0.0F, 0.0F, 0.0F, + 0.0F, 0.0F, 0.0F, 0.0F}); + test.AddInput("B", B_dims, B_data); + + std::vector sequence_lens_dims({batch_size}); + std::vector sequence_lens_data(batch_size, (int)seq_length); + test.AddInput("sequence_lens", sequence_lens_dims, sequence_lens_data); + + std::vector initial_h_dims = {num_directions, batch_size, hidden_size}; + std::vector initial_h_data({0.1F, 0.2F, 0.3F, 0.4F}); + test.AddInput("initial_h", initial_h_dims, initial_h_data); + + std::vector Y_dims = {seq_length, num_directions, batch_size, hidden_size}; + std::vector Y_data({0.98009639F, 0.98009639F, 0.99100745F, 0.99100745F}); + test.AddOutput("Y", Y_dims, Y_data); + + std::vector Y_h_dims{num_directions, batch_size, hidden_size}; + std::vector Y_h_data({0.98009639F, 0.98009639F, 0.99100745F, 0.99100745F}); + test.AddOutput("Y_h", Y_h_dims, Y_h_data); + + test.Run(); +} + typedef enum { RNNOutputY, RNNOutputY_h,