| 969 | } |
| 970 | |
| 971 | void tensor_conv:: |
| 972 | setup( |
| 973 | const tensor& data, |
| 974 | const tensor& filters, |
| 975 | int stride_y_, |
| 976 | int stride_x_, |
| 977 | int padding_y_, |
| 978 | int padding_x_ |
| 979 | ) |
| 980 | { |
| 981 | DLIB_CASSERT(data.k() == filters.k()); |
| 982 | |
| 983 | // if the last call to setup gave the same exact settings then don't do |
| 984 | // anything. |
| 985 | if (data_num_samples == data.num_samples() && |
| 986 | data_k == data.k() && |
| 987 | data_nr == data.nr() && |
| 988 | data_nc == data.nc() && |
| 989 | stride_y_ == stride_y && |
| 990 | stride_x_ == stride_x && |
| 991 | padding_y_ == padding_y && |
| 992 | padding_x_ == padding_x && |
| 993 | filters_num_samples == filters.num_samples() && |
| 994 | filters_k == filters.k() && |
| 995 | filters_nr == filters.nr() && |
| 996 | filters_nc == filters.nc() |
| 997 | ) |
| 998 | { |
| 999 | return; |
| 1000 | } |
| 1001 | |
| 1002 | clear(); |
| 1003 | try |
| 1004 | { |
| 1005 | stride_y = stride_y_; |
| 1006 | stride_x = stride_x_; |
| 1007 | padding_y = padding_y_; |
| 1008 | padding_x = padding_x_; |
| 1009 | data_num_samples = data.num_samples(); |
| 1010 | data_k = data.k(); |
| 1011 | data_nr = data.nr(); |
| 1012 | data_nc = data.nc(); |
| 1013 | filters_num_samples = filters.num_samples(); |
| 1014 | filters_k = filters.k(); |
| 1015 | filters_nr = filters.nr(); |
| 1016 | filters_nc = filters.nc(); |
| 1017 | |
| 1018 | CHECK_CUDNN(cudnnCreateFilterDescriptor((cudnnFilterDescriptor_t*)&filter_handle)); |
| 1019 | CHECK_CUDNN(cudnnSetFilter4dDescriptor((cudnnFilterDescriptor_t)filter_handle, |
| 1020 | CUDNN_DATA_FLOAT, |
| 1021 | CUDNN_TENSOR_NCHW, |
| 1022 | filters.num_samples(), |
| 1023 | filters.k(), |
| 1024 | filters.nr(), |
| 1025 | filters.nc())); |
| 1026 | |
| 1027 | CHECK_CUDNN(cudnnCreateConvolutionDescriptor((cudnnConvolutionDescriptor_t*)&conv_handle)); |
| 1028 | #if CUDNN_MAJOR >= 6 |
nothing calls this directly
no test coverage detected