MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / CheckRNNParameterSize

Function CheckRNNParameterSize

tensorflow/stream_executor/cuda/cuda_dnn.cc:1756–1770  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1754}
1755
1756port::Status CheckRNNParameterSize(
1757 const CudnnHandle& cudnn, const CudnnRnnDescriptor& rnn_desc,
1758 const CudnnRnnSequenceTensorDescriptor& input_desc) {
1759 size_t params_size_in_bytes = 0;
1760 RETURN_IF_CUDNN_ERROR(cudnnGetRNNParamsSize(
1761 /*handle=*/cudnn.handle(), /*rnnDesc=*/rnn_desc.handle(),
1762 /*xDesc=*/input_desc.handles()[0], /*sizeInBytes=*/&params_size_in_bytes,
1763 /*dataType=*/rnn_desc.data_type()));
1764 if (static_cast<int64>(params_size_in_bytes) !=
1765 rnn_desc.ParamsSizeInBytes()) {
1766 return port::Status(port::error::INVALID_ARGUMENT,
1767 "Mismatching RNN parameter size");
1768 }
1769 return port::Status::OK();
1770}
1771
1772port::StatusOr<DeviceMemory<uint8>> CreateRnnWorkspace(
1773 Stream* stream, const CudnnHandle& cudnn,

Callers 2

DoRnnForwardImplMethod · 0.70
DoRnnBackwardImplMethod · 0.70

Calls 5

StatusEnum · 0.50
handleMethod · 0.45
handlesMethod · 0.45
data_typeMethod · 0.45
ParamsSizeInBytesMethod · 0.45

Tested by

no test coverage detected