| 1954 | } |
| 1955 | |
| 1956 | bool CheckRNNParameterSize( |
| 1957 | miopenHandle_t miopen_handle, const MIOpenRnnDescriptor& rnn_desc, |
| 1958 | const MIOpenRnnSequenceTensorDescriptor& input_desc) { |
| 1959 | size_t params_size_in_bytes = 0; |
| 1960 | auto status = wrap::miopenGetRNNParamsSize( |
| 1961 | miopen_handle /*handle*/, rnn_desc.handle() /*rnnDesc*/, |
| 1962 | input_desc.handles()[0] /*xDesc*/, ¶ms_size_in_bytes /*sizeInBytes*/, |
| 1963 | rnn_desc.data_type() /*dataType*/); |
| 1964 | if (status != miopenStatusSuccess) { |
| 1965 | LOG(ERROR) << "Unable to check RNN param size: " << ToString(status); |
| 1966 | return false; |
| 1967 | } |
| 1968 | return static_cast<int64>(params_size_in_bytes) == |
| 1969 | rnn_desc.ParamsSizeInBytes(); |
| 1970 | } |
| 1971 | |
| 1972 | bool CreateRnnWorkspace(Stream* stream, miopenHandle_t miopen_handle, |
| 1973 | const MIOpenRnnDescriptor& rnn_desc, |
no test coverage detected