| 61 | } |
| 62 | |
| 63 | struct MklDnnThreadPool : public threadpool_iface { |
| 64 | MklDnnThreadPool() = default; |
| 65 | |
| 66 | MklDnnThreadPool(OpKernelContext* ctx) |
| 67 | : eigen_interface_(ctx->device() |
| 68 | ->tensorflow_cpu_worker_threads() |
| 69 | ->workers->AsEigenThreadPool()) { |
| 70 | // Set MKL intra thread pool number. |
| 71 | int intra_num = 0; |
| 72 | const char* intra_num_str = getenv("TF_MKL_NUM_INTRAOP"); |
| 73 | const int tf_intra_num = eigen_interface_->NumThreads(); |
| 74 | |
| 75 | if (intra_num_str != NULL) { |
| 76 | intra_num = std::stoi(intra_num_str); |
| 77 | } |
| 78 | intra_num_ = |
| 79 | intra_num > 0 ? std::min(tf_intra_num, intra_num) : tf_intra_num; |
| 80 | dnnl_threadpool_interop_set_max_concurrency(intra_num_); |
| 81 | } |
| 82 | |
| 83 | MklDnnThreadPool(OpKernelContext* ctx, int user_intra_num) |
| 84 | : eigen_interface_(ctx->device() |
| 85 | ->tensorflow_cpu_worker_threads() |
| 86 | ->workers->AsEigenThreadPool()), |
| 87 | intra_num_(user_intra_num) { |
| 88 | // Set MKL intra thread pool number. |
| 89 | int intra_num = 0; |
| 90 | const char* intra_num_str = getenv("TF_MKL_NUM_INTRAOP"); |
| 91 | |
| 92 | if (intra_num_str != NULL) { |
| 93 | intra_num = std::stoi(intra_num_str); |
| 94 | } |
| 95 | intra_num_ = |
| 96 | intra_num > 0 ? std::min(user_intra_num, intra_num) : user_intra_num; |
| 97 | |
| 98 | intra_num_ = intra_num_ > 0 ? intra_num_ : eigen_interface_->NumThreads(); |
| 99 | dnnl_threadpool_interop_set_max_concurrency(intra_num_); |
| 100 | } |
| 101 | |
| 102 | virtual int get_num_threads() const override { |
| 103 | return intra_num_; |
| 104 | } |
| 105 | virtual bool get_in_parallel() const override { |
| 106 | return (eigen_interface_->CurrentThreadId() != -1) ? true : false; |
| 107 | } |
| 108 | virtual uint64_t get_flags() const override { return ASYNCHRONOUS; } |
| 109 | virtual void parallel_for(int n, |
| 110 | const std::function<void(int, int)>& fn) override { |
| 111 | // Should never happen (handled by DNNL) |
| 112 | if (n == 0) return; |
| 113 | |
| 114 | // Should never happen (handled by DNNL) |
| 115 | if (n == 1) { |
| 116 | fn(0, 1); |
| 117 | return; |
| 118 | } |
| 119 | |
| 120 | int nthr = get_num_threads(); |
no outgoing calls
no test coverage detected