| 1894 | |
| 1895 | template <class T> |
| 1896 | bool ExtractAndCheckRnnForward( |
| 1897 | const MIOpenRnnDescriptor& rnn_desc, |
| 1898 | const MIOpenRnnSequenceTensorDescriptor& input_desc, |
| 1899 | const DeviceMemory<T>& input_data, |
| 1900 | const MIOpenRnnStateTensorDescriptor& input_h_desc, |
| 1901 | const DeviceMemory<T>& input_h_data, |
| 1902 | const MIOpenRnnStateTensorDescriptor& input_c_desc, |
| 1903 | const DeviceMemory<T>& input_c_data, const DeviceMemory<T>& params, |
| 1904 | const MIOpenRnnSequenceTensorDescriptor& output_desc, |
| 1905 | const DeviceMemory<T>& output_data, |
| 1906 | const MIOpenRnnStateTensorDescriptor& output_h_desc, |
| 1907 | const DeviceMemory<T>& output_h_data, |
| 1908 | const MIOpenRnnStateTensorDescriptor& output_c_desc, |
| 1909 | const DeviceMemory<T>& output_c_data, RnnModelDims* model_dims) { |
| 1910 | // extract model parameters |
| 1911 | model_dims->num_layers = rnn_desc.num_layers(); |
| 1912 | model_dims->batch_size = input_desc.batch_size(); |
| 1913 | model_dims->seq_length = input_desc.seq_length(); |
| 1914 | model_dims->hidden_size = rnn_desc.hidden_size(); |
| 1915 | model_dims->input_size = input_desc.data_size(); |
| 1916 | model_dims->dir_count = |
| 1917 | (rnn_desc.direction_mode() == miopenRNNbidirection) ? 2 : 1; |
| 1918 | |
| 1919 | // check parameters |
| 1920 | if (!(input_h_desc.num_layers() == |
| 1921 | model_dims->num_layers * model_dims->dir_count && |
| 1922 | input_h_desc.batch_size() == model_dims->batch_size && |
| 1923 | input_h_desc.data_size() == model_dims->hidden_size)) { |
| 1924 | LOG(ERROR) << "Invalid input_h shape"; |
| 1925 | return false; |
| 1926 | } |
| 1927 | if (!(input_h_desc.num_layers() == input_c_desc.num_layers() && |
| 1928 | input_h_desc.batch_size() == input_c_desc.batch_size() && |
| 1929 | input_h_desc.data_size() == input_c_desc.data_size())) { |
| 1930 | LOG(ERROR) << "Invalid input_c shape"; |
| 1931 | return false; |
| 1932 | } |
| 1933 | if (!(output_desc.seq_length() == model_dims->seq_length && |
| 1934 | output_desc.batch_size() == model_dims->batch_size && |
| 1935 | output_desc.data_size() == |
| 1936 | model_dims->hidden_size * model_dims->dir_count)) { |
| 1937 | LOG(ERROR) << "Invalid output shape"; |
| 1938 | return false; |
| 1939 | } |
| 1940 | if (!(input_h_desc.num_layers() == output_h_desc.num_layers() && |
| 1941 | input_h_desc.batch_size() == output_h_desc.batch_size() && |
| 1942 | input_h_desc.data_size() == output_h_desc.data_size())) { |
| 1943 | LOG(ERROR) << "Invalid output_h shape"; |
| 1944 | return false; |
| 1945 | } |
| 1946 | if (!(input_h_desc.num_layers() == output_c_desc.num_layers() && |
| 1947 | input_h_desc.batch_size() == output_c_desc.batch_size() && |
| 1948 | input_h_desc.data_size() == output_c_desc.data_size())) { |
| 1949 | LOG(ERROR) << "Invalid output_h shape"; |
| 1950 | return false; |
| 1951 | } |
| 1952 | |
| 1953 | return true; |
no test coverage detected