MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / do_device_wait_by

Method do_device_wait_by

src/core/impl/comp_node/cpu/comp_node.cpp:1079–1138  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1077
1078#if MGB_HAVE_THREAD
1079void 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);

Callers

nothing calls this directly

Calls 7

locatorMethod · 0.45
cur_recorderMethod · 0.45
loadMethod · 0.45
propertyMethod · 0.45
syncMethod · 0.45
waitMethod · 0.45
add_callbackMethod · 0.45

Tested by

no test coverage detected