已合并
【feat】mstx support push and pop #36207
mei-feiyao创建于 5月20日
【feat】mstx support push and pop #36207
已合并
共 8 个文件变更+326-15
| @@ -9,6 +9,7 @@ class TestMstx(TestCase): | |||
| 9 | range_msg = '' | 9 | range_msg = '' |
| 10 | range_id = 0 | 10 | 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_id | 38 | self.range_id = range_id |
| 38 | self.range_domain = domain | 39 | 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_mark | 66 | torch_npu._C._mstx._mark = stub_mark |
| 41 | torch_npu._C._mstx._mark_on_host = stub_mark_on_host | 67 | torch_npu._C._mstx._mark_on_host = stub_mark_on_host |
| 42 | torch_npu._C._mstx._range_start = stub_range_start | 68 | torch_npu._C._mstx._range_start = stub_range_start |
| 43 | torch_npu._C._mstx._range_start_on_host = stub_range_start_on_host | 69 | torch_npu._C._mstx._range_start_on_host = stub_range_start_on_host |
| 44 | torch_npu._C._mstx._range_end = stub_range_end | 70 | 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 inputs | 76 | # 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 | ||
| 124 | if __name__ == '__main__': | 205 | if __name__ == '__main__': |
| 125 | run_tests() | 206 | run_tests() |
| @@ -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 | +} |
| @@ -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_ERRORS | 138 | END_HANDLE_TH_ERRORS |
| 136 | } | 139 | } |
| @@ -140,20 +143,75 @@ PyObject* THNPModule_mark(PyObject* _unused, PyObject* args) | |||
| 140 | HANDLE_TH_ERRORS | 143 | 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_ERRORS | 158 | 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 | + | ||
| 157 | PyObject* THNPModule_rangeStart(PyObject* _unused, PyObject* args) | 215 | PyObject* THNPModule_rangeStart(PyObject* _unused, PyObject* args) |
| 158 | { | 216 | { |
| 159 | HANDLE_TH_ERRORS | 217 | 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_ERRORS | 234 | 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_ERRORS | 251 | 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_ERRORS | 267 | END_HANDLE_TH_ERRORS |
| 200 | } | 268 | } |
| @@ -202,6 +270,9 @@ PyObject* THNPModule_rangeEnd(PyObject* self, PyObject* args) | |||
| 202 | static std::vector<PyMethodDef> mstxMethods = { | 270 | static 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}, |
| @@ -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 | + | ||
| 41 | MstxMgr::MstxMgr() | 44 | MstxMgr::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 | + | ||
| 66 | int MstxMgr::rangeStart(const char* message, const aclrtStream stream, const char* domain) | 128 | int MstxMgr::rangeStart(const char* message, const aclrtStream stream, const char* domain) |
| 67 | { | 129 | { |
| 68 | if (!isMstxEnable()) { | 130 | if (!isMstxEnable()) { |
| @@ -2,6 +2,7 @@ | |||
| 2 | 2 | ||
| 3 | 3 | ||
| 4 | 4 | ||
| 5 | + | ||
| 5 | 6 | ||
| 6 | 7 | ||
| 7 | 8 | ||
| @@ -20,6 +21,8 @@ class MstxMgr : public torch_npu::toolkit::profiler::Singleton<MstxMgr> { | |||
| 20 | friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>; | 21 | friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>; |
| 21 | public: | 22 | public: |
| 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 | ||
| 50 | private: | 53 | private: |
| 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_; |
| @@ -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 | + | ||
| 152 | inline int mstxRangeStart(const char* message, const aclrtStream stream, const char* domain) | 162 | inline 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); |
| @@ -17,7 +17,7 @@ import functools | |||
| 17 | import torch_npu._C | 17 | import torch_npu._C |
| 18 | from torch_npu.utils.utils import _print_error_log | 18 | from torch_npu.utils.utils import _print_error_log |
| 19 | 19 | ||
| 20 | -__all__ = ["mstx"] | 20 | +__all__ = ["mstx", "annotate"] |
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | def _no_exception_func(default_ret=None): | 23 | def _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 | + | ||
| 61 | + | ||
| 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 | + | ||
| 83 | + | ||
| 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 | 90 | ||
| 61 | 91 | ||
| 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 ret | 131 | return ret |
| 102 | return inner | 132 | return inner |
| 103 | return wrapper | 133 | 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 | + | ||
| 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 | ||
| @@ -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) |