| 1091 | }; |
| 1092 | |
| 1093 | class CudnnRnnParamsDescriptor { |
| 1094 | typedef dnn::RnnDescriptor::ParamsRegions ParamsRegions; |
| 1095 | |
| 1096 | CudnnRnnParamsDescriptor(FilterDescriptor handle, int64 params_size_in_bytes, |
| 1097 | ParamsRegions weights, ParamsRegions biases) |
| 1098 | : handle_(std::move(handle)), |
| 1099 | params_size_in_bytes_(params_size_in_bytes), |
| 1100 | weights_(std::move(weights)), |
| 1101 | biases_(std::move(biases)) {} |
| 1102 | |
| 1103 | public: |
| 1104 | CudnnRnnParamsDescriptor(CudnnRnnParamsDescriptor&&) = default; |
| 1105 | |
| 1106 | static port::StatusOr<CudnnRnnParamsDescriptor> Create( |
| 1107 | const CudnnHandle& cudnn, int input_size, cudnnDataType_t data_type, |
| 1108 | cudnnRNNDescriptor_t rnn_desc, cudnnRNNMode_t rnn_mode, |
| 1109 | cudnnDirectionMode_t direction_mode, int num_layers); |
| 1110 | |
| 1111 | cudnnFilterDescriptor_t handle() const { return handle_.get(); } |
| 1112 | int64 params_size_in_bytes() const { return params_size_in_bytes_; } |
| 1113 | ParamsRegions params_weights() const { return weights_; } |
| 1114 | ParamsRegions params_biases() const { return biases_; } |
| 1115 | |
| 1116 | private: |
| 1117 | FilterDescriptor handle_; |
| 1118 | int64 params_size_in_bytes_; |
| 1119 | ParamsRegions weights_; |
| 1120 | ParamsRegions biases_; |
| 1121 | SE_DISALLOW_COPY_AND_ASSIGN(CudnnRnnParamsDescriptor); |
| 1122 | }; |
| 1123 | |
| 1124 | } // namespace |
| 1125 | |