已合并
【feat】mstx support push and pop #36207
mei-feiyao创建于 5月20日
【feat】mstx support push and pop #36207
已合并
mei-feiyao创建于 5月20日
8 个文件变更+326-15
Mtest/npu/test_mstx.py+81-0
@@ -9,6 +9,7 @@ class TestMstx(TestCase):
9 range_msg = ''9 range_msg = ''
10 range_id = 010 range_id = 0
11 range_domain = ''11 range_domain = ''
12+ push_depth = {}
12 13 
13 def setUp(self):14 def setUp(self):
14 def stub_mark(message: str, stream_id: int = 0, device_index: int = 0,15 def stub_mark(message: str, stream_id: int = 0, device_index: int = 0,
@@ -37,11 +38,39 @@ class TestMstx(TestCase):
37 self.range_id = range_id38 self.range_id = range_id
38 self.range_domain = domain39 self.range_domain = domain
39 40 
41+ def stub_range_push(message: str, stream_id: int = 0, device_index: int = 0,
42+ device_type: int = 0, domain: str = 'default') -> int:
43+ # For simplicity of testing, we use a dict to track the push depth for each domain.
44+ if domain in self.push_depth.keys():
45+ self.push_depth[domain].append(message)
46+ else:
47+ self.push_depth[domain] = [message]
48+ return len(self.push_depth[domain]) - 1
49+ 
50+ def stub_range_push_on_host(message: str, domain: str = 'default') -> int:
51+ if domain in self.push_depth.keys():
52+ self.push_depth[domain].append(message)
53+ else:
54+ self.push_depth[domain] = [message]
55+ return len(self.push_depth[domain]) - 1
56+ 
57+ def stub_range_pop(domain: str = 'default') -> int:
58+ if domain in self.push_depth.keys():
59+ depth = len(self.push_depth[domain]) - 1
60+ if len(self.push_depth[domain]) > 0:
61+ self.push_depth[domain].pop()
62+ return depth
63+ else:
64+ return -1
65+ 
40 torch_npu._C._mstx._mark = stub_mark66 torch_npu._C._mstx._mark = stub_mark
41 torch_npu._C._mstx._mark_on_host = stub_mark_on_host67 torch_npu._C._mstx._mark_on_host = stub_mark_on_host
42 torch_npu._C._mstx._range_start = stub_range_start68 torch_npu._C._mstx._range_start = stub_range_start
43 torch_npu._C._mstx._range_start_on_host = stub_range_start_on_host69 torch_npu._C._mstx._range_start_on_host = stub_range_start_on_host
44 torch_npu._C._mstx._range_end = stub_range_end70 torch_npu._C._mstx._range_end = stub_range_end
71+ torch_npu._C._mstx._range_push = stub_range_push
72+ torch_npu._C._mstx._range_push_on_host = stub_range_push_on_host
73+ torch_npu._C._mstx._range_pop = stub_range_pop
45 74 
46 def test_mark(self):75 def test_mark(self):
47 # invalid inputs76 # invalid inputs
@@ -120,6 +149,58 @@ class TestMstx(TestCase):
120 self.assertEqual(2, self.range_id)149 self.assertEqual(2, self.range_id)
121 self.assertEqual("test_domain1", self.range_domain)150 self.assertEqual("test_domain1", self.range_domain)
122 151 
152+ def test_range_push_will_return_err_value_when_called_with_invalid_inputs(self):
153+ # invalid inputs
154+ ret_id = torch_npu.npu.mstx.range_push("")
155+ self.assertEqual(-1, ret_id)
156+ ret_id = torch_npu.npu.mstx.range_push(message=0)
157+ self.assertEqual(-1, ret_id)
158+ ret_id = torch_npu.npu.mstx.range_push(message="test", stream=None, domain=1)
159+ self.assertEqual(-1, ret_id)
160+ ret_id = torch_npu.npu.mstx.range_push(message="test", stream=1, domain="test")
161+ self.assertEqual(-1, ret_id)
162+ 
163+ def test_range_push_will_return_increased_depth_when_called_with_valid_inputs(self):
164+ # valid inputs
165+ ret_id = torch_npu.npu.mstx.range_push("test1")
166+ self.assertEqual(0, ret_id)
167+ self.assertEqual({"default": ["test1"]}, self.push_depth)
168+ ret_id = torch_npu.npu.mstx.range_push("test2", None, domain="test_domain1")
169+ self.assertEqual(0, ret_id)
170+ self.assertEqual({"default": ["test1"], "test_domain1": ["test2"]}, self.push_depth)
171+ ret_id = torch_npu.npu.mstx.range_push("test3", None, domain="test_domain1")
172+ self.assertEqual(1, ret_id)
173+ self.assertEqual({"default": ["test1"], "test_domain1": ["test2", "test3"]}, self.push_depth)
174+ torch.npu.set_device(0)
175+ current_stream = torch.npu.current_stream()
176+ ret_id = torch_npu.npu.mstx.range_push("test4", current_stream, domain="test_domain2")
177+ self.assertEqual(0, ret_id)
178+ self.assertEqual({"default": ["test1"], "test_domain1": ["test2", "test3"], "test_domain2": ["test4"]}, self.push_depth)
179+ ret_id = torch_npu.npu.mstx.range_push("test5", current_stream, domain="test_domain2")
180+ self.assertEqual(1, ret_id)
181+ self.assertEqual({"default": ["test1"], "test_domain1": ["test2", "test3"], "test_domain2": ["test4", "test5"]}, self.push_depth)
182+ 
183+ def test_range_pop_will_return_err_value_when_called_with_invalid_inputs(self):
184+ # invalid inputs
185+ ret_id = torch_npu.npu.mstx.range_pop(domain=1)
186+ self.assertEqual(-1, ret_id)
187+ ret_id = torch_npu.npu.mstx.range_pop(domain="")
188+ self.assertEqual(-1, ret_id)
189+ 
190+ def test_range_pop_will_return_decreased_depth_when_called_after_range_push_is_called(self):
191+ # valid inputs
192+ self.push_depth = {}
193+ torch_npu.npu.mstx.range_push("test1", None, domain="test_domain1")
194+ torch_npu.npu.mstx.range_push("test2", None, domain="test_domain1")
195+ ret_id = torch_npu.npu.mstx.range_pop(domain="test_domain1")
196+ self.assertEqual(1, ret_id)
197+ self.assertEqual({"test_domain1": ["test1"]}, self.push_depth)
198+ ret_id = torch_npu.npu.mstx.range_pop(domain="test_domain1")
199+ self.assertEqual(0, ret_id)
200+ self.assertEqual({"test_domain1": []}, self.push_depth)
201+ ret_id = torch_npu.npu.mstx.range_pop(domain="test_domain1")
202+ self.assertEqual(-1, ret_id)
203+ self.assertEqual({"test_domain1": []}, self.push_depth)
123 204 
124if __name__ == '__main__':205if __name__ == '__main__':
125 run_tests()206 run_tests()
Mtest/torch_npu_schema.json+19-1
@@ -1058,6 +1058,12 @@
1058 "torch_npu.npu.mstx.mark": {1058 "torch_npu.npu.mstx.mark": {
1059 "signature": "(message: str, stream=None, domain: str = 'default')"1059 "signature": "(message: str, stream=None, domain: str = 'default')"
1060 },1060 },
1061+ "torch_npu.npu.mstx.range_push": {
1062+ "signature": "(message: str, stream=None, domain: str = 'default') -> int"
1063+ },
1064+ "torch_npu.npu.mstx.range_pop": {
1065+ "signature": "(domain: str = 'default') -> int"
1066+ },
1061 "torch_npu.npu.mstx.range_start": {1067 "torch_npu.npu.mstx.range_start": {
1062 "signature": "(message: str, stream=None, domain: str = 'default') -> int"1068 "signature": "(message: str, stream=None, domain: str = 'default') -> int"
1063 },1069 },
@@ -1067,6 +1073,9 @@
1067 "torch_npu.npu.mstx.mstx_range": {1073 "torch_npu.npu.mstx.mstx_range": {
1068 "signature": "(message: str, stream=None, domain: str = 'default')"1074 "signature": "(message: str, stream=None, domain: str = 'default')"
1069 },1075 },
1076+ "torch_npu.npu.mstx.annotate": {
1077+ "signature": "(message: str = '', stream=None, domain: str = 'default')"
1078+ },
1070 "torch_npu.npu.reset_accumulated_memory_stats": {1079 "torch_npu.npu.reset_accumulated_memory_stats": {
1071 "signature": "(device=None)"1080 "signature": "(device=None)"
1072 },1081 },
@@ -1355,6 +1364,15 @@
1355 "torch_npu.npu.mstx.mstx.mstx_range": {1364 "torch_npu.npu.mstx.mstx.mstx_range": {
1356 "signature": "(message: str, stream=None, domain: str = 'default')"1365 "signature": "(message: str, stream=None, domain: str = 'default')"
1357 },1366 },
1367+ "torch_npu.npu.mstx.mstx.range_push": {
1368+ "signature": "(message: str, stream=None, domain: str = 'default') -> int"
1369+ },
1370+ "torch_npu.npu.mstx.mstx.annotate": {
1371+ "signature": "(message: str = '', stream=None, domain: str = 'default')"
1372+ },
1373+ "torch_npu.npu.mstx.mstx.range_pop": {
1374+ "signature": "(domain: str = 'default') -> int"
1375+ },
1358 "torch_npu.npu.npu_config.finalize_dump": {1376 "torch_npu.npu.npu_config.finalize_dump": {
1359 "signature": "()"1377 "signature": "()"
1360 },1378 },
@@ -2886,4 +2904,4 @@
2886 "signature": "(std::vector<std::string>& op_type, std::vector<at::Tensor>& tensors, std::vector<int64_t> remote_rank_list) -> c10::intrusive_ptr<c10d::Work>",2904 "signature": "(std::vector<std::string>& op_type, std::vector<at::Tensor>& tensors, std::vector<int64_t> remote_rank_list) -> c10::intrusive_ptr<c10d::Work>",
2887 "file": "torch_npu/csrc/distributed/ProcessGroupHCCL.hpp"2905 "file": "torch_npu/csrc/distributed/ProcessGroupHCCL.hpp"
2888 }2906 }
2889-}2907+}
Mtorch_npu/csrc/profiler/init.cpp+83-12
@@ -130,7 +130,10 @@ PyObject* THNPModule_markOnHost(PyObject* _unused, PyObject* args)
130 if (!PyArg_ParseTuple(args, "ss", &message, &domain)) {130 if (!PyArg_ParseTuple(args, "ss", &message, &domain)) {
131 return nullptr;131 return nullptr;
132 }132 }
133- mstxMark(message, nullptr, domain);133+ {
134+ pybind11::gil_scoped_release no_gil;
135+ mstxMark(message, nullptr, domain);
136+ }
134 Py_RETURN_NONE;137 Py_RETURN_NONE;
135 END_HANDLE_TH_ERRORS138 END_HANDLE_TH_ERRORS
136}139}
@@ -140,20 +143,75 @@ PyObject* THNPModule_mark(PyObject* _unused, PyObject* args)
140 HANDLE_TH_ERRORS143 HANDLE_TH_ERRORS
141 const char* message;144 const char* message;
142 const char* domain;145 const char* domain;
143- PyObject* stream_o = nullptr;
144 int64_t stream_id = 0;146 int64_t stream_id = 0;
145 int64_t device_index = 0;147 int64_t device_index = 0;
146 int64_t device_type = 0;148 int64_t device_type = 0;
147 if (!PyArg_ParseTuple(args, "sLLLs", &message, &stream_id, &device_index, &device_type, &domain)) {149 if (!PyArg_ParseTuple(args, "sLLLs", &message, &stream_id, &device_index, &device_type, &domain)) {
148 return nullptr;150 return nullptr;
149 }151 }
150- auto stream = c10_npu::NPUStream::unpack3(152+ {
151- stream_id, device_index, static_cast<c10::DeviceType>(device_type));153+ pybind11::gil_scoped_release no_gil;
152- mstxMark(message, stream.stream(false), domain);154+ auto stream = c10_npu::NPUStream::unpack3(stream_id, device_index, static_cast<c10::DeviceType>(device_type));
155+ mstxMark(message, stream.stream(false), domain);
156+ }
153 Py_RETURN_NONE;157 Py_RETURN_NONE;
154 END_HANDLE_TH_ERRORS158 END_HANDLE_TH_ERRORS
155}159}
156 160 
161+PyObject* THNPModule_rangePushOnHost(PyObject* _unused, PyObject* args)
162+{
163+ HANDLE_TH_ERRORS
164+ const char* message = nullptr;
165+ const char* domain = nullptr;
166+ if (!PyArg_ParseTuple(args, "ss", &message, &domain)) {
167+ return nullptr;
168+ }
169+ int id;
170+ {
171+ pybind11::gil_scoped_release no_gil;
172+ id = mstxRangePush(message, nullptr, domain);
173+ }
174+ return PyLong_FromLong(static_cast<long>(id));
175+ END_HANDLE_TH_ERRORS
176+}
177+ 
178+PyObject* THNPModule_rangePush(PyObject* _unused, PyObject* args)
179+{
180+ HANDLE_TH_ERRORS
181+ const char* message = nullptr;
182+ const char* domain = nullptr;
183+ int64_t stream_id = 0;
184+ int64_t device_index = 0;
185+ int64_t device_type = 0;
186+ if (!PyArg_ParseTuple(args, "sLLLs", &message, &stream_id, &device_index, &device_type, &domain)) {
187+ return nullptr;
188+ }
189+ int id;
190+ {
191+ pybind11::gil_scoped_release no_gil;
192+ auto stream = c10_npu::NPUStream::unpack3(stream_id, device_index, static_cast<c10::DeviceType>(device_type));
193+ id = mstxRangePush(message, stream.stream(false), domain);
194+ }
195+ return PyLong_FromLong(static_cast<long>(id));
196+ END_HANDLE_TH_ERRORS
197+}
198+ 
199+PyObject* THNPModule_rangePop(PyObject* _unused, PyObject* args)
200+{
201+ HANDLE_TH_ERRORS
202+ const char* domain = nullptr;
203+ if (!PyArg_ParseTuple(args, "s", &domain)) {
204+ return nullptr;
205+ }
206+ int id;
207+ {
208+ pybind11::gil_scoped_release no_gil;
209+ id = mstxRangePop(domain);
210+ }
211+ return PyLong_FromLong(static_cast<long>(id));
212+ END_HANDLE_TH_ERRORS
213+}
214+ 
157PyObject* THNPModule_rangeStart(PyObject* _unused, PyObject* args)215PyObject* THNPModule_rangeStart(PyObject* _unused, PyObject* args)
158{216{
159 HANDLE_TH_ERRORS217 HANDLE_TH_ERRORS
@@ -166,10 +224,13 @@ PyObject* THNPModule_rangeStart(PyObject* _unused, PyObject* args)
166 if (!PyArg_ParseTuple(args, "sLLLs", &message, &stream_id, &device_index, &device_type, &domain)) {224 if (!PyArg_ParseTuple(args, "sLLLs", &message, &stream_id, &device_index, &device_type, &domain)) {
167 return nullptr;225 return nullptr;
168 }226 }
169- auto stream = c10_npu::NPUStream::unpack3(227+ int id;
170- stream_id, device_index, static_cast<c10::DeviceType>(device_type));228+ {
171- int id = mstxRangeStart(message, stream.stream(false), domain);229+ pybind11::gil_scoped_release no_gil;
172- return PyLong_FromLong(id);230+ auto stream = c10_npu::NPUStream::unpack3(stream_id, device_index, static_cast<c10::DeviceType>(device_type));
231+ id = mstxRangeStart(message, stream.stream(false), domain);
232+ }
233+ return PyLong_FromLong(static_cast<long>(id));
173 END_HANDLE_TH_ERRORS234 END_HANDLE_TH_ERRORS
174}235}
175 236 
@@ -181,8 +242,12 @@ PyObject* THNPModule_rangeStartOnHost(PyObject* _unused, PyObject* args)
181 if (!PyArg_ParseTuple(args, "ss", &message, &domain)) {242 if (!PyArg_ParseTuple(args, "ss", &message, &domain)) {
182 return nullptr;243 return nullptr;
183 }244 }
184- int id = mstxRangeStart(message, nullptr, domain);245+ int id;
185- return PyLong_FromLong(id);246+ {
247+ pybind11::gil_scoped_release no_gil;
248+ id = mstxRangeStart(message, nullptr, domain);
249+ }
250+ return PyLong_FromLong(static_cast<long>(id));
186 END_HANDLE_TH_ERRORS251 END_HANDLE_TH_ERRORS
187}252}
188 253 
@@ -194,7 +259,10 @@ PyObject* THNPModule_rangeEnd(PyObject* self, PyObject* args)
194 if (!PyArg_ParseTuple(args, "is", &rangeId, &domain)) {259 if (!PyArg_ParseTuple(args, "is", &rangeId, &domain)) {
195 return nullptr;260 return nullptr;
196 }261 }
197- mstxRangeEnd(rangeId, domain);262+ {
263+ pybind11::gil_scoped_release no_gil;
264+ mstxRangeEnd(rangeId, domain);
265+ }
198 Py_RETURN_NONE;266 Py_RETURN_NONE;
199 END_HANDLE_TH_ERRORS267 END_HANDLE_TH_ERRORS
200}268}
@@ -202,6 +270,9 @@ PyObject* THNPModule_rangeEnd(PyObject* self, PyObject* args)
202static std::vector<PyMethodDef> mstxMethods = {270static std::vector<PyMethodDef> mstxMethods = {
203 {"_mark_on_host", (PyCFunction)THNPModule_markOnHost, METH_VARARGS, nullptr},271 {"_mark_on_host", (PyCFunction)THNPModule_markOnHost, METH_VARARGS, nullptr},
204 {"_mark", (PyCFunction)THNPModule_mark, METH_VARARGS, nullptr},272 {"_mark", (PyCFunction)THNPModule_mark, METH_VARARGS, nullptr},
273+ {"_range_push_on_host", (PyCFunction)THNPModule_rangePushOnHost, METH_VARARGS, nullptr},
274+ {"_range_push", (PyCFunction)THNPModule_rangePush, METH_VARARGS, nullptr},
275+ {"_range_pop", (PyCFunction)THNPModule_rangePop, METH_VARARGS, nullptr},
205 {"_range_start_on_host", (PyCFunction)THNPModule_rangeStartOnHost, METH_VARARGS, nullptr},276 {"_range_start_on_host", (PyCFunction)THNPModule_rangeStartOnHost, METH_VARARGS, nullptr},
206 {"_range_start", (PyCFunction)THNPModule_rangeStart, METH_VARARGS, nullptr},277 {"_range_start", (PyCFunction)THNPModule_rangeStart, METH_VARARGS, nullptr},
207 {"_range_end", (PyCFunction)THNPModule_rangeEnd, METH_VARARGS, nullptr},278 {"_range_end", (PyCFunction)THNPModule_rangeEnd, METH_VARARGS, nullptr},
Mtorch_npu/csrc/profiler/mstx_mgr.cpp+62-0
@@ -38,6 +38,9 @@ void rangeEndImpl(int ptRangeId, mstxDomainHandle_t domain)
38 }38 }
39}39}
40 40 
41+thread_local std::unordered_map<std::string, std::stack<int>> MstxMgr::domainPushDepthStacks_ = {};
42+thread_local bool MstxMgr::pushWithStream_ = false;
43+ 
41MstxMgr::MstxMgr()44MstxMgr::MstxMgr()
42{45{
43}46}
@@ -63,6 +66,65 @@ void MstxMgr::mark(const char* message, const aclrtStream stream, const char* do
63 at_npu::native::OpCommand::RunOpApiV2("mstx_mark_op", mark_call);66 at_npu::native::OpCommand::RunOpApiV2("mstx_mark_op", mark_call);
64}67}
65 68 
69+int MstxMgr::rangePush(const char* message, const aclrtStream stream, const char* domain)
70+{
71+ if (!isMstxEnable()) {
72+ return -1;
73+ }
74+ std::string domainStr(domain);
75+ if (!isMstxTxDomainEnable(domainStr)) {
76+ return -1;
77+ }
78+ // depth before push to return
79+ int ret = domainPushDepthStacks_[domainStr].size();
80+ int id = ptRangeId_++;
81+ domainPushDepthStacks_[domainStr].push(id);
82+ mstxDomainHandle_t domainHandle = createProfDomain(domainStr);
83+ // currently use rangeStartImpl to implement rangePushImpl.
84+ if (stream == nullptr) {
85+ pushWithStream_ = false;
86+ rangeStartImpl(message, nullptr, id, domainHandle);
87+ return ret;
88+ }
89+ pushWithStream_ = true;
90+ auto range_push_call = [msg_ptr = std::make_shared<std::string>(message), stream, id, domainHandle]() -> int {
91+ rangeStartImpl(msg_ptr->c_str(), stream, id, domainHandle);
92+ return 0;
93+ };
94+ at_npu::native::OpCommand::RunOpApiV2("mstx_range_push_op", range_push_call);
95+ return ret;
96+}
97+ 
98+int MstxMgr::rangePop(const char* domain)
99+{
100+ if (!isMstxEnable()) {
101+ return -1;
102+ }
103+ std::string domainStr(domain);
104+ if (domainPushDepthStacks_[domainStr].empty()) {
105+ return -1;
106+ }
107+ if (!isMstxTxDomainEnable(domainStr)) {
108+ return -1;
109+ }
110+ int id = domainPushDepthStacks_[domainStr].top();
111+ domainPushDepthStacks_[domainStr].pop();
112+ // depth after pop to return
113+ int ret = domainPushDepthStacks_[domainStr].size();
114+ mstxDomainHandle_t domainHandle = createProfDomain(domainStr);
115+ // currently use rangeEndImpl to implement rangePopImpl.
116+ if (!pushWithStream_) {
117+ rangeEndImpl(id, domainHandle);
118+ return ret;
119+ }
120+ auto range_pop_call = [domainHandle, id]() -> int {
121+ rangeEndImpl(id, domainHandle);
122+ return 0;
123+ };
124+ at_npu::native::OpCommand::RunOpApiV2("mstx_range_pop_op", range_pop_call);
125+ return ret;
126+}
127+ 
66int MstxMgr::rangeStart(const char* message, const aclrtStream stream, const char* domain)128int MstxMgr::rangeStart(const char* message, const aclrtStream stream, const char* domain)
67{129{
68 if (!isMstxEnable()) {130 if (!isMstxEnable()) {
Mtorch_npu/csrc/profiler/mstx_mgr.h+5-0
@@ -2,6 +2,7 @@
2 2 
3#include <atomic>3#include <atomic>
4#include <mutex>4#include <mutex>
5+#include <stack>
5#include <unordered_map>6#include <unordered_map>
6#include <unordered_set>7#include <unordered_set>
7#include "torch_npu/csrc/framework/interface/MstxInterface.h"8#include "torch_npu/csrc/framework/interface/MstxInterface.h"
@@ -20,6 +21,8 @@ class MstxMgr : public torch_npu::toolkit::profiler::Singleton<MstxMgr> {
20friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>;21friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>;
21public:22public:
22 void mark(const char* message, const aclrtStream stream, const char* domain);23 void mark(const char* message, const aclrtStream stream, const char* domain);
24+ int rangePush(const char* message, const aclrtStream stream, const char* domain);
25+ int rangePop(const char* domain);
23 int rangeStart(const char* message, const aclrtStream stream, const char* domain);26 int rangeStart(const char* message, const aclrtStream stream, const char* domain);
24 void rangeEnd(int ptRangeId, const char* domain);27 void rangeEnd(int ptRangeId, const char* domain);
25 28 
@@ -48,6 +51,8 @@ private:
48 bool isMsptiTxEnableImpl();51 bool isMsptiTxEnableImpl();
49 52 
50private:53private:
54+ static thread_local std::unordered_map<std::string, std::stack<int>> domainPushDepthStacks_;
55+ static thread_local bool pushWithStream_;
51 std::atomic<int> ptRangeId_{1};56 std::atomic<int> ptRangeId_{1};
52 std::unordered_set<int> ptRangeIdsWithStream_;57 std::unordered_set<int> ptRangeIdsWithStream_;
53 std::mutex mtx_;58 std::mutex mtx_;
Mtorch_npu/csrc/profiler/npu_profiler.h+10-0
@@ -149,6 +149,16 @@ inline void mstxMark(const char* message, const aclrtStream stream, const char*
149 }149 }
150}150}
151 151 
152+inline int mstxRangePush(const char* message, const aclrtStream stream, const char* domain)
153+{
154+ return MstxMgr::GetInstance()->rangePush(message, stream, domain);
155+}
156+ 
157+inline int mstxRangePop(const char* domain)
158+{
159+ return MstxMgr::GetInstance()->rangePop(domain);
160+}
161+ 
152inline int mstxRangeStart(const char* message, const aclrtStream stream, const char* domain)162inline int mstxRangeStart(const char* message, const aclrtStream stream, const char* domain)
153{163{
154 return MstxMgr::GetInstance()->rangeStart(message, stream, domain);164 return MstxMgr::GetInstance()->rangeStart(message, stream, domain);
Mtorch_npu/npu/mstx.py+65-1
@@ -17,7 +17,7 @@ import functools
17import torch_npu._C17import torch_npu._C
18from torch_npu.utils.utils import _print_error_log18from torch_npu.utils.utils import _print_error_log
19 19 
20-__all__ = ["mstx"]20+__all__ = ["mstx", "annotate"]
21 21 
22 22 
23def _no_exception_func(default_ret=None):23def _no_exception_func(default_ret=None):
@@ -57,6 +57,36 @@ class mstx:
57 else:57 else:
58 torch_npu._C._mstx._mark_on_host(message, domain)58 torch_npu._C._mstx._mark_on_host(message, domain)
59 59 
60+ @staticmethod
61+ @_no_exception_func()
62+ def range_push(message: str, stream=None, domain: str = 'default') -> int:
63+ if not message or not isinstance(message, str):
64+ warnings.warn("Invalid message for mstx.range_push func. Please input valid message string.")
65+ return -1
66+ if not domain or not isinstance(domain, str):
67+ warnings.warn("Invalid domain for mstx.range_push func. Please input valid domain string.")
68+ return -1
69+ if stream:
70+ if isinstance(stream, torch_npu.npu.streams.Stream):
71+ return torch_npu._C._mstx._range_push(message,
72+ stream.stream_id,
73+ stream.device_index,
74+ stream.device_type,
75+ domain)
76+ else:
77+ warnings.warn("Invalid stream for mstx.range_push func. Please input valid stream.")
78+ return -1
79+ else:
80+ return torch_npu._C._mstx._range_push_on_host(message, domain)
81+ 
82+ @staticmethod
83+ @_no_exception_func()
84+ def range_pop(domain: str = 'default') -> int:
85+ if not domain or not isinstance(domain, str):
86+ warnings.warn("Invalid domain for mstx.range_pop func. Please input valid domain string.")
87+ return -1
88+ return torch_npu._C._mstx._range_pop(domain)
89+ 
60 @staticmethod90 @staticmethod
61 @_no_exception_func()91 @_no_exception_func()
62 def range_start(message: str, stream=None, domain: str = 'default') -> int:92 def range_start(message: str, stream=None, domain: str = 'default') -> int:
@@ -101,3 +131,37 @@ class mstx:
101 return ret131 return ret
102 return inner132 return inner
103 return wrapper133 return wrapper
134+ 
135+ 
136+class annotate:
137+ def __init__(self, message: str = '', stream=None, domain: str = 'default'):
138+ self.message = message
139+ self.stream = stream
140+ self.domain = domain
141+ self.range_id = None
142+ 
143+ def __enter__(self):
144+ self.range_id = mstx.range_start(self.message, self.stream, self.domain)
145+ return self
146+ 
147+ def __exit__(self, exc_type, exc_val, exc_tb):
148+ if self.range_id is not None:
149+ mstx.range_end(self.range_id, self.domain)
150+ self.range_id = None
151+ 
152+ def __call__(self, func):
153+ if not self.message:
154+ self.message = func.__name__
155+ 
156+ @functools.wraps(func)
157+ def inner(*args, **kwargs):
158+ range_id = mstx.range_start(self.message, self.stream, self.domain)
159+ try:
160+ result = func(*args, **kwargs)
161+ finally:
162+ mstx.range_end(range_id, self.domain)
163+ return result
164+ 
165+ return inner
166+ 
167+mstx.annotate = annotate
Mtorch_npu/profiler/analysis/prof_parse/_fwk_file_parser.py+1-1
@@ -249,7 +249,7 @@ class FwkFileParser:
249 None if not torch_op.args.get(Constant.INPUT_SHAPES) else str2id_manager.get_id_from_str(torch_op.args.get(Constant.INPUT_SHAPES)),249 None if not torch_op.args.get(Constant.INPUT_SHAPES) else str2id_manager.get_id_from_str(torch_op.args.get(Constant.INPUT_SHAPES)),
250 None if not torch_op.args.get(Constant.CALL_STACK) else call_chain_id_manager.get_callchain_id_from_callstack(torch_op.args.get(Constant.CALL_STACK)),250 None if not torch_op.args.get(Constant.CALL_STACK) else call_chain_id_manager.get_callchain_id_from_callstack(torch_op.args.get(Constant.CALL_STACK)),
251 ApiType.TORCH_OP]251 ApiType.TORCH_OP]
252- if torch_op.name == "mstx_mark_op":252+ if torch_op.name.startswith("mstx_"):
253 mstx_mark_apis.append(api)253 mstx_mark_apis.append(api)
254 else:254 else:
255 torch_op_apis.append(api)255 torch_op_apis.append(api)