| 80 | |
| 81 | template <typename T> |
| 82 | void MklPoolingFwdPrimitive<T>::Execute(const T* src_data, T* dst_data, |
| 83 | void* ws_data, |
| 84 | std::shared_ptr<stream> fwd_stream) { |
| 85 | #ifdef DNNL_AARCH64_USE_ACL |
| 86 | mutex_lock lock(primitive_execution_mu_); |
| 87 | #endif |
| 88 | #ifdef ENABLE_DNNL_THREADPOOL |
| 89 | context_.src_mem->set_data_handle( |
| 90 | static_cast<void*>(const_cast<T*>(src_data)), *fwd_stream); |
| 91 | context_.dst_mem->set_data_handle(static_cast<void*>(dst_data), *fwd_stream); |
| 92 | if (context_.alg_kind == ALGORITHM::pooling_max && |
| 93 | context_.prop_kind == |
| 94 | prop_kind::forward_training) { // Max pooling must have workspace. |
| 95 | DCHECK(ws_data != nullptr); |
| 96 | context_.ws_mem->set_data_handle(ws_data, *fwd_stream); |
| 97 | } |
| 98 | #else |
| 99 | context_.src_mem->set_data_handle( |
| 100 | static_cast<void*>(const_cast<T*>(src_data))); |
| 101 | context_.dst_mem->set_data_handle(static_cast<void*>(dst_data)); |
| 102 | if (context_.alg_kind == ALGORITHM::pooling_max && |
| 103 | context_.prop_kind == |
| 104 | prop_kind::forward_training) { // Max pooling must have workspace. |
| 105 | DCHECK(ws_data != nullptr); |
| 106 | context_.ws_mem->set_data_handle(ws_data); |
| 107 | } |
| 108 | #endif // ENABLE_DNNL_THREADPOOL |
| 109 | execute_primitives(context_.fwd_primitives, fwd_stream, context_.net_args); |
| 110 | |
| 111 | // Set back data handle. |
| 112 | context_.src_mem->set_data_handle(DummyData); |
| 113 | context_.dst_mem->set_data_handle(DummyData); |
| 114 | if (context_.alg_kind == ALGORITHM::pooling_max && |
| 115 | context_.prop_kind == |
| 116 | prop_kind::forward_training) { // Max pooling must have workspace. |
| 117 | DCHECK(ws_data != nullptr); |
| 118 | context_.ws_mem->set_data_handle(DummyData); |
| 119 | } |
| 120 | } |
| 121 | |
| 122 | template class MklPoolingFwdPrimitive<float>; |
| 123 | template class MklPoolingFwdPrimitive<quint8>; |