TODO(mirkov): add another test which directly compares to TF once TOCO supports the conversion from dynamic_rnn with BasicRNNCell.
| 816 | // TODO(mirkov): add another test which directly compares to TF once TOCO |
| 817 | // supports the conversion from dynamic_rnn with BasicRNNCell. |
| 818 | TEST(BidirectionalRNNOpTest, BlackBoxTest) { |
| 819 | BidirectionalRNNOpModel rnn(/*batches=*/2, /*sequence_len=*/16, |
| 820 | /*fw_units=*/16, /*bw_units=*/16, |
| 821 | /*input_size=*/8, /*aux_input_size=*/0, |
| 822 | /*aux_input_mode=*/AuxInputMode::kNoAuxInput, |
| 823 | /*time_major=*/false, |
| 824 | /*merge_outputs=*/false); |
| 825 | rnn.SetFwWeights(weights); |
| 826 | rnn.SetBwWeights(weights); |
| 827 | rnn.SetFwBias(biases); |
| 828 | rnn.SetBwBias(biases); |
| 829 | rnn.SetFwRecurrentWeights(recurrent_weights); |
| 830 | rnn.SetBwRecurrentWeights(recurrent_weights); |
| 831 | |
| 832 | const int input_sequence_size = rnn.input_size() * rnn.sequence_len(); |
| 833 | float* batch_start = rnn_input; |
| 834 | float* batch_end = batch_start + input_sequence_size; |
| 835 | rnn.SetInput(0, batch_start, batch_end); |
| 836 | rnn.SetInput(input_sequence_size, batch_start, batch_end); |
| 837 | |
| 838 | rnn.Invoke(); |
| 839 | |
| 840 | float* golden_fw_start = rnn_golden_fw_output; |
| 841 | float* golden_fw_end = |
| 842 | golden_fw_start + rnn.num_fw_units() * rnn.sequence_len(); |
| 843 | std::vector<float> fw_expected; |
| 844 | fw_expected.insert(fw_expected.end(), golden_fw_start, golden_fw_end); |
| 845 | fw_expected.insert(fw_expected.end(), golden_fw_start, golden_fw_end); |
| 846 | EXPECT_THAT(rnn.GetFwOutput(), ElementsAreArray(ArrayFloatNear(fw_expected))); |
| 847 | |
| 848 | float* golden_bw_start = rnn_golden_bw_output; |
| 849 | float* golden_bw_end = |
| 850 | golden_bw_start + rnn.num_bw_units() * rnn.sequence_len(); |
| 851 | std::vector<float> bw_expected; |
| 852 | bw_expected.insert(bw_expected.end(), golden_bw_start, golden_bw_end); |
| 853 | bw_expected.insert(bw_expected.end(), golden_bw_start, golden_bw_end); |
| 854 | EXPECT_THAT(rnn.GetBwOutput(), ElementsAreArray(ArrayFloatNear(bw_expected))); |
| 855 | } |
| 856 | |
| 857 | // Same as BlackBox test, but input is reshuffled to time_major format. |
| 858 | TEST(BidirectionalRNNOpTest, BlackBoxTestTimeMajor) { |
nothing calls this directly
no test coverage detected