| 1077 | |
| 1078 | #if MGB_HAVE_THREAD |
| 1079 | void CpuCompNode::CpuDispatchableBase::EventImpl::do_device_wait_by(Impl* cn_impl) { |
| 1080 | { |
| 1081 | auto locator = m_comp_node_impl->locator(); |
| 1082 | if (locator.device == Locator::DEVICE_CPU_DEFAULT && |
| 1083 | !static_cast<CpuCompNode::CompNodeRecorderImpl*>(m_comp_node_impl) |
| 1084 | ->cur_recorder()) { |
| 1085 | auto v0 = m_record_nr_req.load(std::memory_order_relaxed), |
| 1086 | v1 = m_record_nr_finish.load(std::memory_order_relaxed); |
| 1087 | mgb_assert( |
| 1088 | v0 && v0 == v1, |
| 1089 | "event on cpu:default hasn't been recorded inplace."); |
| 1090 | return; |
| 1091 | } |
| 1092 | } |
| 1093 | |
| 1094 | { |
| 1095 | auto type = cn_impl->env().property().type; |
| 1096 | mgb_throw_if( |
| 1097 | type != CompNode::DeviceType::CPU && |
| 1098 | type != CompNode::DeviceType::CUDA |
| 1099 | && type != CompNode::DeviceType::ATLAS && |
| 1100 | type != CompNode::DeviceType::CAMBRICON |
| 1101 | , |
| 1102 | MegBrainError, |
| 1103 | "currently CPU can only wait for CPU, CUDA, ATLAS" |
| 1104 | ); |
| 1105 | } |
| 1106 | |
| 1107 | if (cn_impl->env().property().type == CompNode::DeviceType::ATLAS) { |
| 1108 | #if MGB_ATLAS |
| 1109 | return m_comp_node_impl->sync(); |
| 1110 | #else |
| 1111 | mgb_throw(MegBrainError, "Atlas comp_node used but ATLAS BUILD not enabled"); |
| 1112 | #endif |
| 1113 | } else if (cn_impl->env().property().type == CompNode::DeviceType::CAMBRICON) { |
| 1114 | #if MGB_CAMBRICON |
| 1115 | return m_comp_node_impl->sync(); |
| 1116 | #else |
| 1117 | mgb_throw( |
| 1118 | MegBrainError, |
| 1119 | "Cambricon comp_node used but CAMBRICON BUILD not enabled"); |
| 1120 | #endif |
| 1121 | } |
| 1122 | |
| 1123 | auto version = m_record_nr_req.load(std::memory_order_relaxed); |
| 1124 | mgb_assert(version, "device wait on non-recorded event"); |
| 1125 | |
| 1126 | auto waiter = [this, version]() { |
| 1127 | while (m_record_nr_finish.load(std::memory_order_acquire) < version) { |
| 1128 | std::unique_lock<std::mutex> lk{m_dev_wait_mtx}; |
| 1129 | if (m_record_nr_finish.load(std::memory_order_acquire) >= version) { |
| 1130 | break; |
| 1131 | } |
| 1132 | m_dev_wait_cv.wait(lk); |
| 1133 | } |
| 1134 | m_dev_wait_nr_waiter.fetch_sub(1, std::memory_order_release); |
| 1135 | }; |
| 1136 | m_dev_wait_nr_waiter.fetch_add(1, std::memory_order_release); |
nothing calls this directly
no test coverage detected