已合并
fix: 修复dflow非0卡部署context校验失败并保证API层context透明 #4697
lining23666创建于 29 天前
fix: 修复dflow非0卡部署context校验失败并保证API层context透明 #4697
已合并
共 2 个文件变更+49-0
| @@ -12,6 +12,7 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 15 | 16 | ||
| 16 | namespace ge { | 17 | namespace ge { |
| 17 | namespace { | 18 | namespace { |
| @@ -89,6 +90,8 @@ Status ProxyEventManager::GetProxyPid(int32_t device_id, int32_t &pid) { | |||
| 89 | 90 | ||
| 90 | Status ProxyEventManager::SubmitEventSync(int32_t device_id, uint32_t sub_event_id, char_t *msg, size_t msg_len, | 91 | Status ProxyEventManager::SubmitEventSync(int32_t device_id, uint32_t sub_event_id, char_t *msg, size_t msg_len, |
| 91 | rtEschedEventReply_t *ack) { | 92 | rtEschedEventReply_t *ack) { |
| 93 | + // rt esched api depends on thread-private runtime context, ensure current thread binds the device first | ||
| 94 | + GE_CHK_STATUS_RET(RtsApiUtils::SetDevice(device_id), "Failed to set device, device_id = %d.", device_id); | ||
| 92 | int32_t proxy_pid = -1; | 95 | int32_t proxy_pid = -1; |
| 93 | GE_CHK_STATUS_RET(GetProxyPid(device_id, proxy_pid), "Failed to get proxy pid, device_id = %d.", device_id); | 96 | GE_CHK_STATUS_RET(GetProxyPid(device_id, proxy_pid), "Failed to get proxy pid, device_id = %d.", device_id); |
| 94 | uint32_t group_id = 0U; | 97 | uint32_t group_id = 0U; |
| @@ -44,6 +44,36 @@ std::atomic_bool g_dflow_ge_initialized{false}; | |||
| 44 | std::mutex g_dflow_ge_release_mutex; // DFlowInitialize, DFlowFinalize and ~DFlowSession use | 44 | std::mutex g_dflow_ge_release_mutex; // DFlowInitialize, DFlowFinalize and ~DFlowSession use |
| 45 | std::shared_ptr<DFlowSessionManager> g_dflow_session_manager; | 45 | std::shared_ptr<DFlowSessionManager> g_dflow_session_manager; |
| 46 | 46 | ||
| 47 | +// Save caller's runtime context on entry and restore it on exit, so that internal device/context | ||
| 48 | +// switching of dflow is transparent to api callers. No-op when the caller thread has no context. | ||
| 49 | +// Skip restoring when the context is unchanged during the call, in case it was destroyed midway | ||
| 50 | +// (e.g. by internal aclrtResetDevice) and restoring an invalid handle would fail. | ||
| 51 | +class RtCtxGuard { | ||
| 52 | + public: | ||
| 53 | + RtCtxGuard() { | ||
| 54 | + if (aclrtGetCurrentContext(&old_ctx_) != ACL_SUCCESS) { | ||
| 55 | + old_ctx_ = nullptr; | ||
| 56 | + } | ||
| 57 | + } | ||
| 58 | + ~RtCtxGuard() { | ||
| 59 | + if (old_ctx_ == nullptr) { | ||
| 60 | + return; | ||
| 61 | + } | ||
| 62 | + aclrtContext cur_ctx = nullptr; | ||
| 63 | + if ((aclrtGetCurrentContext(&cur_ctx) != ACL_SUCCESS) || (cur_ctx == old_ctx_)) { | ||
| 64 | + return; | ||
| 65 | + } | ||
| 66 | + if (aclrtSetCurrentContext(old_ctx_) != ACL_SUCCESS) { | ||
| 67 | + GELOGW("Failed to restore runtime context after dflow api call, the context may have been destroyed."); | ||
| 68 | + } | ||
| 69 | + } | ||
| 70 | + RtCtxGuard(const RtCtxGuard &) = delete; | ||
| 71 | + RtCtxGuard &operator=(const RtCtxGuard &) = delete; | ||
| 72 | + | ||
| 73 | + private: | ||
| 74 | + aclrtContext old_ctx_ = nullptr; | ||
| 75 | +}; | ||
| 76 | + | ||
| 47 | void DFlowFinalizeImpl() { | 77 | void DFlowFinalizeImpl() { |
| 48 | GELOGT(TRACE_INIT, "DFlowFinalize start."); | 78 | GELOGT(TRACE_INIT, "DFlowFinalize start."); |
| 49 | if (g_dflow_session_manager != nullptr) { | 79 | if (g_dflow_session_manager != nullptr) { |
| @@ -62,6 +92,7 @@ void DFlowFinalizeImpl() { | |||
| 62 | } // namespace | 92 | } // namespace |
| 63 | 93 | ||
| 64 | Status DFlowInitialize(const std::map<AscendString, AscendString> &options) { | 94 | Status DFlowInitialize(const std::map<AscendString, AscendString> &options) { |
| 95 | + RtCtxGuard ctx_guard; | ||
| 65 | std::lock_guard<std::mutex> lock(g_dflow_ge_release_mutex); | 96 | std::lock_guard<std::mutex> lock(g_dflow_ge_release_mutex); |
| 66 | if (g_dflow_ge_initialized) { | 97 | if (g_dflow_ge_initialized) { |
| 67 | GELOGW("DFlowInitialize is called more than once"); | 98 | GELOGW("DFlowInitialize is called more than once"); |
| @@ -101,6 +132,7 @@ Status DFlowInitialize(const std::map<AscendString, AscendString> &options) { | |||
| 101 | 132 | ||
| 102 | // DFlow finalize, releasing all resources | 133 | // DFlow finalize, releasing all resources |
| 103 | Status DFlowFinalize() { | 134 | Status DFlowFinalize() { |
| 135 | + RtCtxGuard ctx_guard; | ||
| 104 | GRAPH_PROFILING_REG(gert::GeProfInfoType::kGEFinalize); | 136 | GRAPH_PROFILING_REG(gert::GeProfInfoType::kGEFinalize); |
| 105 | std::lock_guard<std::mutex> lock(g_dflow_ge_release_mutex); | 137 | std::lock_guard<std::mutex> lock(g_dflow_ge_release_mutex); |
| 106 | DFlowFinalizeImpl(); | 138 | DFlowFinalizeImpl(); |
| @@ -136,12 +168,14 @@ void ConstructSession(const std::map<std::string, std::string> &options, Session | |||
| 136 | } // namespace | 168 | } // namespace |
| 137 | 169 | ||
| 138 | DFlowSession::DFlowSession(const std::map<AscendString, AscendString> &options) { | 170 | DFlowSession::DFlowSession(const std::map<AscendString, AscendString> &options) { |
| 171 | + RtCtxGuard ctx_guard; | ||
| 139 | std::map<std::string, std::string> str_options; | 172 | std::map<std::string, std::string> str_options; |
| 140 | ConvertAscendStringMap(options, str_options); | 173 | ConvertAscendStringMap(options, str_options); |
| 141 | ConstructSession(str_options, dflow_session_impl_); | 174 | ConstructSession(str_options, dflow_session_impl_); |
| 142 | } | 175 | } |
| 143 | 176 | ||
| 144 | DFlowSession::~DFlowSession() { | 177 | DFlowSession::~DFlowSession() { |
| 178 | + RtCtxGuard ctx_guard; | ||
| 145 | if (dflow_session_impl_ == nullptr) { | 179 | if (dflow_session_impl_ == nullptr) { |
| 146 | return; | 180 | return; |
| 147 | } | 181 | } |
| @@ -175,6 +209,7 @@ DFlowSession::~DFlowSession() { | |||
| 175 | 209 | ||
| 176 | Status DFlowSession::AddGraph(uint32_t graph_id, const FlowGraph &graph, | 210 | Status DFlowSession::AddGraph(uint32_t graph_id, const FlowGraph &graph, |
| 177 | const std::map<AscendString, AscendString> &options) { | 211 | const std::map<AscendString, AscendString> &options) { |
| 212 | + RtCtxGuard ctx_guard; | ||
| 178 | if (!g_dflow_ge_initialized) { | 213 | if (!g_dflow_ge_initialized) { |
| 179 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Construct][DFlowSession]Failed because GEInitialize was not called before."); | 214 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Construct][DFlowSession]Failed because GEInitialize was not called before."); |
| 180 | REPORT_INNER_ERR_MSG("E19999", "Creating session failed because GEInitialize was not called before."); | 215 | REPORT_INNER_ERR_MSG("E19999", "Creating session failed because GEInitialize was not called before."); |
| @@ -203,6 +238,7 @@ Status DFlowSession::AddGraph(uint32_t graph_id, const FlowGraph &graph, | |||
| 203 | } | 238 | } |
| 204 | 239 | ||
| 205 | Status DFlowSession::RemoveGraph(uint32_t graph_id) { | 240 | Status DFlowSession::RemoveGraph(uint32_t graph_id) { |
| 241 | + RtCtxGuard ctx_guard; | ||
| 206 | if (!g_dflow_ge_initialized) { | 242 | if (!g_dflow_ge_initialized) { |
| 207 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Construct][DFlowSession]Failed because GEInitialize was not called before."); | 243 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Construct][DFlowSession]Failed because GEInitialize was not called before."); |
| 208 | REPORT_INNER_ERR_MSG("E19999", "Creating session failed because GEInitialize was not called before."); | 244 | REPORT_INNER_ERR_MSG("E19999", "Creating session failed because GEInitialize was not called before."); |
| @@ -224,6 +260,7 @@ Status DFlowSession::RemoveGraph(uint32_t graph_id) { | |||
| 224 | } | 260 | } |
| 225 | 261 | ||
| 226 | Status DFlowSession::BuildGraph(uint32_t graph_id, const std::vector<ge::Tensor> &inputs) { | 262 | Status DFlowSession::BuildGraph(uint32_t graph_id, const std::vector<ge::Tensor> &inputs) { |
| 263 | + RtCtxGuard ctx_guard; | ||
| 227 | if (!g_dflow_ge_initialized) { | 264 | if (!g_dflow_ge_initialized) { |
| 228 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Construct][DFlowSession]Failed because GEInitialize was not called before."); | 265 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Construct][DFlowSession]Failed because GEInitialize was not called before."); |
| 229 | REPORT_INNER_ERR_MSG("E19999", "Creating session failed because GEInitialize was not called before."); | 266 | REPORT_INNER_ERR_MSG("E19999", "Creating session failed because GEInitialize was not called before."); |
| @@ -253,11 +290,13 @@ uint64_t DFlowSession::GetSessionId() const { | |||
| 253 | 290 | ||
| 254 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<Tensor> &inputs, const DataFlowInfo &info, | 291 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<Tensor> &inputs, const DataFlowInfo &info, |
| 255 | int32_t timeout) { | 292 | int32_t timeout) { |
| 293 | + RtCtxGuard ctx_guard; | ||
| 256 | return FeedDataFlowGraph(graph_id, {}, inputs, info, timeout); | 294 | return FeedDataFlowGraph(graph_id, {}, inputs, info, timeout); |
| 257 | } | 295 | } |
| 258 | 296 | ||
| 259 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, | 297 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, |
| 260 | const std::vector<Tensor> &inputs, const DataFlowInfo &info, int32_t timeout) { | 298 | const std::vector<Tensor> &inputs, const DataFlowInfo &info, int32_t timeout) { |
| 299 | + RtCtxGuard ctx_guard; | ||
| 261 | if (!g_dflow_ge_initialized) { | 300 | if (!g_dflow_ge_initialized) { |
| 262 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Feed][Data]Failed because GEInitialize was not called before."); | 301 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Feed][Data]Failed because GEInitialize was not called before."); |
| 263 | REPORT_INNER_ERR_MSG("E19999", "Feed data failed because GEInitialize was not called before."); | 302 | REPORT_INNER_ERR_MSG("E19999", "Feed data failed because GEInitialize was not called before."); |
| @@ -279,11 +318,13 @@ Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint | |||
| 279 | } | 318 | } |
| 280 | 319 | ||
| 281 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<FlowMsgPtr> &inputs, int32_t timeout) { | 320 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<FlowMsgPtr> &inputs, int32_t timeout) { |
| 321 | + RtCtxGuard ctx_guard; | ||
| 282 | return FeedDataFlowGraph(graph_id, {}, inputs, timeout); | 322 | return FeedDataFlowGraph(graph_id, {}, inputs, timeout); |
| 283 | } | 323 | } |
| 284 | 324 | ||
| 285 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, | 325 | Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, |
| 286 | const std::vector<FlowMsgPtr> &inputs, int32_t timeout) { | 326 | const std::vector<FlowMsgPtr> &inputs, int32_t timeout) { |
| 327 | + RtCtxGuard ctx_guard; | ||
| 287 | GE_CHK_BOOL_RET_STATUS(g_dflow_ge_initialized, FAILED, | 328 | GE_CHK_BOOL_RET_STATUS(g_dflow_ge_initialized, FAILED, |
| 288 | "[Feed][FlowMsg]Failed because GEInitialize was not called before."); | 329 | "[Feed][FlowMsg]Failed because GEInitialize was not called before."); |
| 289 | 330 | ||
| @@ -301,6 +342,7 @@ Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint | |||
| 301 | 342 | ||
| 302 | Status DFlowSession::FeedRawData(uint32_t graph_id, const std::vector<RawData> &raw_data_list, uint32_t index, | 343 | Status DFlowSession::FeedRawData(uint32_t graph_id, const std::vector<RawData> &raw_data_list, uint32_t index, |
| 303 | const DataFlowInfo &info, int32_t timeout) { | 344 | const DataFlowInfo &info, int32_t timeout) { |
| 345 | + RtCtxGuard ctx_guard; | ||
| 304 | if (!g_dflow_ge_initialized) { | 346 | if (!g_dflow_ge_initialized) { |
| 305 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Feed][RawData]Failed because GEInitialize was not called before."); | 347 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Feed][RawData]Failed because GEInitialize was not called before."); |
| 306 | REPORT_INNER_ERR_MSG("E19999", "Feed raw data failed because GEInitialize was not called before."); | 348 | REPORT_INNER_ERR_MSG("E19999", "Feed raw data failed because GEInitialize was not called before."); |
| @@ -322,11 +364,13 @@ Status DFlowSession::FeedRawData(uint32_t graph_id, const std::vector<RawData> & | |||
| 322 | 364 | ||
| 323 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, std::vector<Tensor> &outputs, DataFlowInfo &info, | 365 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, std::vector<Tensor> &outputs, DataFlowInfo &info, |
| 324 | int32_t timeout) { | 366 | int32_t timeout) { |
| 367 | + RtCtxGuard ctx_guard; | ||
| 325 | return FetchDataFlowGraph(graph_id, {}, outputs, info, timeout); | 368 | return FetchDataFlowGraph(graph_id, {}, outputs, info, timeout); |
| 326 | } | 369 | } |
| 327 | 370 | ||
| 328 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, | 371 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, |
| 329 | std::vector<Tensor> &outputs, DataFlowInfo &info, int32_t timeout) { | 372 | std::vector<Tensor> &outputs, DataFlowInfo &info, int32_t timeout) { |
| 373 | + RtCtxGuard ctx_guard; | ||
| 330 | if (!g_dflow_ge_initialized) { | 374 | if (!g_dflow_ge_initialized) { |
| 331 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Fetch][Data]Failed because GEInitialize was not called before."); | 375 | GELOGE(GE_CLI_GE_NOT_INITIALIZED, "[Fetch][Data]Failed because GEInitialize was not called before."); |
| 332 | REPORT_INNER_ERR_MSG("E19999", "Fetch data failed because GEInitialize was not called before."); | 376 | REPORT_INNER_ERR_MSG("E19999", "Fetch data failed because GEInitialize was not called before."); |
| @@ -350,11 +394,13 @@ Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uin | |||
| 350 | } | 394 | } |
| 351 | 395 | ||
| 352 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, std::vector<FlowMsgPtr> &outputs, int32_t timeout) { | 396 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, std::vector<FlowMsgPtr> &outputs, int32_t timeout) { |
| 397 | + RtCtxGuard ctx_guard; | ||
| 353 | return FetchDataFlowGraph(graph_id, {}, outputs, timeout); | 398 | return FetchDataFlowGraph(graph_id, {}, outputs, timeout); |
| 354 | } | 399 | } |
| 355 | 400 | ||
| 356 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, | 401 | Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes, |
| 357 | std::vector<FlowMsgPtr> &outputs, int32_t timeout) { | 402 | std::vector<FlowMsgPtr> &outputs, int32_t timeout) { |
| 403 | + RtCtxGuard ctx_guard; | ||
| 358 | GE_CHK_BOOL_RET_STATUS(g_dflow_ge_initialized, FAILED, | 404 | GE_CHK_BOOL_RET_STATUS(g_dflow_ge_initialized, FAILED, |
| 359 | "[Fetch][FlowMsg]Failed because GEInitialize was not called before."); | 405 | "[Fetch][FlowMsg]Failed because GEInitialize was not called before."); |
| 360 | GE_CHECK_NOTNULL(dflow_session_impl_); | 406 | GE_CHECK_NOTNULL(dflow_session_impl_); |