已合并
fix: 修复dflow非0卡部署context校验失败并保证API层context透明 #4697
lining23666创建于 29 天前
fix: 修复dflow非0卡部署context校验失败并保证API层context透明 #4697
已合并
lining23666创建于 29 天前
共 2 个文件变更+49-0
@@ -12,6 +12,7 @@
12#include "framework/common/debug/log.h"12#include "framework/common/debug/log.h"
13#include "framework/common/ge_inner_error_codes.h"13#include "framework/common/ge_inner_error_codes.h"
14#include "mmpa/mmpa_api.h"14#include "mmpa/mmpa_api.h"
15+#include "common/utils/rts_api_utils.h"
15 16 
16namespace ge {17namespace ge {
17namespace {18namespace {
@@ -89,6 +90,8 @@ Status ProxyEventManager::GetProxyPid(int32_t device_id, int32_t &pid) {
89 90 
90Status ProxyEventManager::SubmitEventSync(int32_t device_id, uint32_t sub_event_id, char_t *msg, size_t msg_len,91Status 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};
44std::mutex g_dflow_ge_release_mutex; // DFlowInitialize, DFlowFinalize and ~DFlowSession use44std::mutex g_dflow_ge_release_mutex; // DFlowInitialize, DFlowFinalize and ~DFlowSession use
45std::shared_ptr<DFlowSessionManager> g_dflow_session_manager;45std::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+ 
47void DFlowFinalizeImpl() {77void 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} // namespace92} // namespace
63 93 
64Status DFlowInitialize(const std::map<AscendString, AscendString> &options) {94Status 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 resources133// DFlow finalize, releasing all resources
103Status DFlowFinalize() {134Status 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} // namespace168} // namespace
137 169 
138DFlowSession::DFlowSession(const std::map<AscendString, AscendString> &options) {170DFlowSession::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 
144DFlowSession::~DFlowSession() {177DFlowSession::~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 
176Status DFlowSession::AddGraph(uint32_t graph_id, const FlowGraph &graph,210Status 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 
205Status DFlowSession::RemoveGraph(uint32_t graph_id) {240Status 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 
226Status DFlowSession::BuildGraph(uint32_t graph_id, const std::vector<ge::Tensor> &inputs) {262Status 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 
254Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<Tensor> &inputs, const DataFlowInfo &info,291Status 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 
259Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes,297Status 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 
281Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<FlowMsgPtr> &inputs, int32_t timeout) {320Status 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 
285Status DFlowSession::FeedDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes,325Status 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 
302Status DFlowSession::FeedRawData(uint32_t graph_id, const std::vector<RawData> &raw_data_list, uint32_t index,343Status 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 
323Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, std::vector<Tensor> &outputs, DataFlowInfo &info,365Status 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 
328Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes,371Status 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 
352Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, std::vector<FlowMsgPtr> &outputs, int32_t timeout) {396Status 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 
356Status DFlowSession::FetchDataFlowGraph(uint32_t graph_id, const std::vector<uint32_t> &indexes,401Status 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_);