已合并
Add mstx dotting to obtain PTA memory pool data #18242
czy6创建于 2025年2月23日
Add mstx dotting to obtain PTA memory pool data #18242
已合并
czy6创建于 2025年2月23日
refs/pull/18242/head合入到master
8 个文件变更+293-36
@@ -13,7 +13,61 @@ typedef uint64_t mstxRangeId;
13 13 
14struct mstxDomainRegistration_st;14struct mstxDomainRegistration_st;
15typedef struct mstxDomainRegistration_st mstxDomainRegistration_t;15typedef struct mstxDomainRegistration_st mstxDomainRegistration_t;
16-typedef mstxDomainRegistration_t* mstxDomainhandle_t;16+typedef mstxDomainRegistration_t* mstxDomainHandle_t;
17+ 
18+struct mstxMemHeap_st;
19+typedef struct mstxMemHeap_st mstxMemHeap_t;
20+typedef mstxMemHeap_t* mstxMemHeapHandle_t;
21+ 
22+struct mstxMemRegion_st;
23+typedef struct mstxMemRegion_st mstxMemRegion_t;
24+typedef mstxMemRegion_t* mstxMemRegionHandle_t;
25+ 
26+typedef struct mstxMemVirtualRangeDesc_t {
27+ uint32_t deviceId;
28+ const void* ptr;
29+ uint64_t size;
30+} mstxMemVirtualRangeDesc_t;
31+ 
32+typedef enum mstxMemHeapUsageType {
33+ MSTX_MEM_HEAP_USAGE_TYPE_SUB_ALLOCATOR = 0,
34+} mstxMemHeapUsageType;
35+ 
36+typedef enum mstxMemType {
37+ MSTX_MEM_TYPE_VIRTUAL_ADDRESS = 0,
38+} mstxMemType;
39+ 
40+typedef struct mstxMemHeapDesc_t {
41+ mstxMemHeapUsageType usage;
42+ mstxMemType type;
43+ const void* typeSpecificDesc;
44+} mstxMemHeapDesc_t;
45+ 
46+typedef struct mstxMemRegionsRegisterBatch_t {
47+ mstxMemHeapHandle_t heap;
48+ mstxMemType regionType;
49+ size_t regionCount;
50+ const void* regionDescArray;
51+ mstxMemRegionHandle_t* regionHandleArrayOut;
52+} mstxMemRegionsRegisterBatch_t;
53+ 
54+typedef enum mstxMemRegionRefType {
55+ MSTX_MEM_REGION_REF_TYPE_POINTER = 0,
56+ MSTX_MEM_REGION_REF_TYPE_HANDLE
57+} mstxMemRegionRefType;
58+ 
59+typedef struct mstxMemRegionRef_t {
60+ mstxMemRegionRefType refType;
61+ union {
62+ const void* pointer;
63+ mstxMemRegionHandle_t handle;
64+ };
65+} mstxMemRegionRef_t;
66+ 
67+typedef struct mstxMemRegionsUnregisterBatch_t {
68+ size_t refCount;
69+ const mstxMemRegionRef_t* refArray;
70+} mstxMemRegionsUnregisterBatch_t;
17 71 
18ACL_FUNC_VISIBILITY void mstxMarkA(const char* message, aclrtStream stream);72ACL_FUNC_VISIBILITY void mstxMarkA(const char* message, aclrtStream stream);
19 73 
@@ -21,16 +75,24 @@ ACL_FUNC_VISIBILITY mstxRangeId mstxRangeStartA(const char* message, aclrtStream
21 75 
22ACL_FUNC_VISIBILITY void mstxRangeEnd(mstxRangeId id);76ACL_FUNC_VISIBILITY void mstxRangeEnd(mstxRangeId id);
23 77 
24-ACL_FUNC_VISIBILITY mstxDomainhandle_t mstxDomainCreateA(const char* name);78+ACL_FUNC_VISIBILITY mstxDomainHandle_t mstxDomainCreateA(const char* name);
25 79 
26-ACL_FUNC_VISIBILITY void mstxDomainDestroy(mstxDomainhandle_t handle);80+ACL_FUNC_VISIBILITY void mstxDomainDestroy(mstxDomainHandle_t handle);
27 81 
28-ACL_FUNC_VISIBILITY void mstxDomainMarkA(mstxDomainhandle_t handle, const char* message, aclrtStream stream);82+ACL_FUNC_VISIBILITY void mstxDomainMarkA(mstxDomainHandle_t handle, const char* message, aclrtStream stream);
29 83 
30-ACL_FUNC_VISIBILITY mstxRangeId mstxDomainRangeStartA(mstxDomainhandle_t handle, const char* message,84+ACL_FUNC_VISIBILITY mstxRangeId mstxDomainRangeStartA(mstxDomainHandle_t handle, const char* message,
czy6
czy6czy62025年2月28日

同版本修改,不存在兼容性问题

likedislike
31 aclrtStream stream);85 aclrtStream stream);
32 86 
33-ACL_FUNC_VISIBILITY void mstxDomainRangeEnd(mstxDomainhandle_t handle, mstxRangeId id);87+ACL_FUNC_VISIBILITY void mstxDomainRangeEnd(mstxDomainHandle_t handle, mstxRangeId id);
88+ 
89+ACL_FUNC_VISIBILITY mstxMemHeapHandle_t mstxMemHeapRegister(mstxDomainHandle_t domain, const mstxMemHeapDesc_t* desc);
90+ 
91+ACL_FUNC_VISIBILITY void mstxMemHeapUnregister(mstxDomainHandle_t domain, mstxMemHeapHandle_t heap);
92+ 
93+ACL_FUNC_VISIBILITY void mstxMemRegionsRegister(mstxDomainHandle_t domain, const mstxMemRegionsRegisterBatch_t* desc);
94+ 
95+ACL_FUNC_VISIBILITY void mstxMemRegionsUnregister(mstxDomainHandle_t domain, const mstxMemRegionsUnregisterBatch_t* desc);
34 96 
35#ifdef __cplusplus97#ifdef __cplusplus
36}98}
@@ -1261,6 +1261,9 @@ class DeviceCachingAllocator {
1261 stats.allocated_bytes[static_cast<size_t>(StatType::AGGREGATE)].current);1261 stats.allocated_bytes[static_cast<size_t>(StatType::AGGREGATE)].current);
1262 1262 
1263#ifndef BUILD_LIBTORCH1263#ifndef BUILD_LIBTORCH
1264+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
1265+ mstxMemVirtualRangeDesc_t desc{block->device, block->ptr, block->size};
1266+ torch_npu::profiler::MstxMgr::GetInstance()->memRegionsRegister(msleaksDomain, &desc);
1264 torch_npu::profiler::reportMemoryDataToNpuProfiler({1267 torch_npu::profiler::reportMemoryDataToNpuProfiler({
1265 static_cast<int8_t>(c10::DeviceType::PrivateUse1),1268 static_cast<int8_t>(c10::DeviceType::PrivateUse1),
1266 block->device,1269 block->device,
@@ -1320,6 +1323,8 @@ class DeviceCachingAllocator {
1320 stats.reserved_bytes[static_cast<size_t>(StatType::AGGREGATE)].current,1323 stats.reserved_bytes[static_cast<size_t>(StatType::AGGREGATE)].current,
1321 stats.allocated_bytes[static_cast<size_t>(StatType::AGGREGATE)].current);1324 stats.allocated_bytes[static_cast<size_t>(StatType::AGGREGATE)].current);
1322#ifndef BUILD_LIBTORCH1325#ifndef BUILD_LIBTORCH
1326+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
1327+ torch_npu::profiler::MstxMgr::GetInstance()->memRegionsUnregister(msleaksDomain, orig_block_ptr);
1323 torch_npu::profiler::reportMemoryDataToNpuProfiler({1328 torch_npu::profiler::reportMemoryDataToNpuProfiler({
1324 static_cast<int8_t>(c10::DeviceType::PrivateUse1),1329 static_cast<int8_t>(c10::DeviceType::PrivateUse1),
1325 block->device,1330 block->device,
@@ -1669,7 +1674,11 @@ class DeviceCachingAllocator {
1669 for_each_selected_stat_type(stat_types, [&](size_t stat_type) {1674 for_each_selected_stat_type(stat_types, [&](size_t stat_type) {
1670 update_stat(stats.reserved_bytes[stat_type], mapped_range.size);1675 update_stat(stats.reserved_bytes[stat_type], mapped_range.size);
1671 });1676 });
1672- 1677+#ifndef BUILD_LIBTORCH
1678+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
1679+ mstxMemVirtualRangeDesc_t desc{to_map->device, mapped_range.ptr, mapped_range.size};
1680+ torch_npu::profiler::MstxMgr::GetInstance()->memHeapRegister(msleaksDomain, &desc);
1681+#endif
1673 record_trace(1682 record_trace(
1674 TraceEntry::SEGMENT_MAP,1683 TraceEntry::SEGMENT_MAP,
1675 int64_t(mapped_range.ptr),1684 int64_t(mapped_range.ptr),
@@ -2048,6 +2057,11 @@ class DeviceCachingAllocator {
2048 2057 
2049 // p.block came from new, not cudaMalloc. It should not be nullptr here.2058 // p.block came from new, not cudaMalloc. It should not be nullptr here.
2050 TORCH_INTERNAL_ASSERT(p.block != nullptr && p.block->ptr != nullptr);2059 TORCH_INTERNAL_ASSERT(p.block != nullptr && p.block->ptr != nullptr);
2060+#ifndef BUILD_LIBTORCH
2061+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
2062+ mstxMemVirtualRangeDesc_t desc{p.block->device, p.block->ptr, p.block->size};
2063+ torch_npu::profiler::MstxMgr::GetInstance()->memHeapRegister(msleaksDomain, &desc);
2064+#endif
2051 record_trace(2065 record_trace(
2052 TraceEntry::SEGMENT_ALLOC,2066 TraceEntry::SEGMENT_ALLOC,
2053 int64_t(p.block->ptr),2067 int64_t(p.block->ptr),
@@ -2165,7 +2179,10 @@ class DeviceCachingAllocator {
2165 2179 
2166 if (block->size >= CachingAllocatorConfig::max_split_size())2180 if (block->size >= CachingAllocatorConfig::max_split_size())
2167 update_stat(stats.oversize_segments, -1);2181 update_stat(stats.oversize_segments, -1);
2168- 2182+#ifndef BUILD_LIBTORCH
2183+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
2184+ torch_npu::profiler::MstxMgr::GetInstance()->memHeapUnregister(msleaksDomain, block->ptr);
2185+#endif
2169 ASCEND_LOGD("pta_memory acl_free: free_size = %zu", block->size);2186 ASCEND_LOGD("pta_memory acl_free: free_size = %zu", block->size);
2170 2187 
2171 pool->blocks.erase(block);2188 pool->blocks.erase(block);
@@ -2223,7 +2240,10 @@ class DeviceCachingAllocator {
2223 for_each_selected_stat_type(stat_types, [&](size_t stat_type) {2240 for_each_selected_stat_type(stat_types, [&](size_t stat_type) {
2224 update_stat(stats.reserved_bytes[stat_type], -unmapped.size);2241 update_stat(stats.reserved_bytes[stat_type], -unmapped.size);
2225 });2242 });
2226- 2243+#ifndef BUILD_LIBTORCH
2244+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
2245+ torch_npu::profiler::MstxMgr::GetInstance()->memHeapUnregister(msleaksDomain, block->ptr);
2246+#endif
2227 record_trace(2247 record_trace(
2228 TraceEntry::SEGMENT_UNMAP,2248 TraceEntry::SEGMENT_UNMAP,
2229 int64_t(unmapped.ptr),2249 int64_t(unmapped.ptr),
@@ -53,6 +53,8 @@ public:
53 NPU_CHECK_ERROR(c10_npu::acl::AclrtSynchronizeDeviceWithTimeout());53 NPU_CHECK_ERROR(c10_npu::acl::AclrtSynchronizeDeviceWithTimeout());
54 NPU_CHECK_ERROR(aclrtFree(block->data_ptr));54 NPU_CHECK_ERROR(aclrtFree(block->data_ptr));
55#ifndef BUILD_LIBTORCH55#ifndef BUILD_LIBTORCH
56+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
57+ torch_npu::profiler::MstxMgr::GetInstance()->memRegionsUnregister(msleaksDomain, block->data_ptr);
56 record_mem_size_decrement(block->size);58 record_mem_size_decrement(block->size);
57 const c10_npu::impl::PyCallbackTrigger* trigger = c10_npu::impl::NPUTrace::getTrace();59 const c10_npu::impl::PyCallbackTrigger* trigger = c10_npu::impl::NPUTrace::getTrace();
58 if (C10_UNLIKELY(trigger)) {60 if (C10_UNLIKELY(trigger)) {
@@ -89,6 +91,9 @@ public:
89 91 
90 ASCEND_LOGD("NPUWorkspaceAllocator malloc by AclrtMallocAlign32: size=%zu", block->size);92 ASCEND_LOGD("NPUWorkspaceAllocator malloc by AclrtMallocAlign32: size=%zu", block->size);
91#ifndef BUILD_LIBTORCH93#ifndef BUILD_LIBTORCH
94+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
95+ mstxMemVirtualRangeDesc_t desc{device, block->data_ptr, block->size};
96+ torch_npu::profiler::MstxMgr::GetInstance()->memRegionsRegister(msleaksDomain, &desc);
92 record_mem_size_increment(block->size);97 record_mem_size_increment(block->size);
93 torch_npu::profiler::reportMemoryDataToNpuProfiler({98 torch_npu::profiler::reportMemoryDataToNpuProfiler({
94 static_cast<int8_t>(c10::DeviceType::PrivateUse1),99 static_cast<int8_t>(c10::DeviceType::PrivateUse1),
@@ -131,6 +136,8 @@ public:
131 ASCEND_LOGI("NPUWorkspaceAllocator free by aclrtFree: size=%zu", block_pair.second->size);136 ASCEND_LOGI("NPUWorkspaceAllocator free by aclrtFree: size=%zu", block_pair.second->size);
132 NPU_CHECK_ERROR(aclrtFree(block_pair.second->data_ptr));137 NPU_CHECK_ERROR(aclrtFree(block_pair.second->data_ptr));
133#ifndef BUILD_LIBTORCH138#ifndef BUILD_LIBTORCH
139+ mstxDomainHandle_t msleaksDomain = torch_npu::profiler::MstxMgr::GetInstance()->createDomain(torch_npu::profiler::DOMAIN_MSLEAKS.c_str());
140+ torch_npu::profiler::MstxMgr::GetInstance()->memRegionsUnregister(msleaksDomain, block_pair.second->data_ptr);
134 record_mem_size_decrement(block_pair.second->size);141 record_mem_size_decrement(block_pair.second->size);
135 const c10_npu::impl::PyCallbackTrigger* trigger = c10_npu::impl::NPUTrace::getTrace();142 const c10_npu::impl::PyCallbackTrigger* trigger = c10_npu::impl::NPUTrace::getTrace();
136 if (C10_UNLIKELY(trigger)) {143 if (C10_UNLIKELY(trigger)) {
@@ -24,6 +24,10 @@ LOAD_FUNCTION(mstxDomainDestroy)
24LOAD_FUNCTION(mstxDomainMarkA)24LOAD_FUNCTION(mstxDomainMarkA)
25LOAD_FUNCTION(mstxDomainRangeStartA)25LOAD_FUNCTION(mstxDomainRangeStartA)
26LOAD_FUNCTION(mstxDomainRangeEnd)26LOAD_FUNCTION(mstxDomainRangeEnd)
27+LOAD_FUNCTION(mstxMemHeapRegister)
28+LOAD_FUNCTION(mstxMemHeapUnregister)
29+LOAD_FUNCTION(mstxMemRegionsRegister)
30+LOAD_FUNCTION(mstxMemRegionsUnregister)
27 31 
28// save python range id with cann mstx range id.32// save python range id with cann mstx range id.
29// when mstx.range_end(id) is called, we can check if this id is invalid33// when mstx.range_end(id) is called, we can check if this id is invalid
@@ -128,9 +132,9 @@ void MstxRangeEnd(int ptRangeId)
128 g_rangeIdMap.erase(iter);132 g_rangeIdMap.erase(iter);
129}133}
130 134 
131-mstxDomainhandle_t MstxDomainCreateA(const char* name)135+mstxDomainHandle_t MstxDomainCreateA(const char* name)
132{136{
133- using MstxDomainCreateAFunc = mstxDomainhandle_t (*)(const char*);137+ using MstxDomainCreateAFunc = mstxDomainHandle_t (*)(const char*);
134 static MstxDomainCreateAFunc func = nullptr;138 static MstxDomainCreateAFunc func = nullptr;
135 static bool noFuncFlag = false;139 static bool noFuncFlag = false;
136 if (noFuncFlag) {140 if (noFuncFlag) {
@@ -147,9 +151,9 @@ mstxDomainhandle_t MstxDomainCreateA(const char* name)
147 return func(name);151 return func(name);
148}152}
149 153 
150-void MstxDomainDestroy(mstxDomainhandle_t handle)154+void MstxDomainDestroy(mstxDomainHandle_t handle)
151{155{
152- using MstxDomainDestroyFunc = void (*)(mstxDomainhandle_t);156+ using MstxDomainDestroyFunc = void (*)(mstxDomainHandle_t);
153 static MstxDomainDestroyFunc func = nullptr;157 static MstxDomainDestroyFunc func = nullptr;
154 static bool noFuncFlag = false;158 static bool noFuncFlag = false;
155 if (noFuncFlag) {159 if (noFuncFlag) {
@@ -166,9 +170,9 @@ void MstxDomainDestroy(mstxDomainhandle_t handle)
166 func(handle);170 func(handle);
167}171}
168 172 
169-void MstxDomainMarkA(mstxDomainhandle_t handle, const char* message, aclrtStream stream)173+void MstxDomainMarkA(mstxDomainHandle_t handle, const char* message, aclrtStream stream)
170{174{
171- using MstxDomainMarkAFunc = void (*)(mstxDomainhandle_t, const char*, aclrtStream);175+ using MstxDomainMarkAFunc = void (*)(mstxDomainHandle_t, const char*, aclrtStream);
172 static MstxDomainMarkAFunc func = nullptr;176 static MstxDomainMarkAFunc func = nullptr;
173 static bool noFuncFlag = false;177 static bool noFuncFlag = false;
174 if (noFuncFlag) {178 if (noFuncFlag) {
@@ -185,9 +189,9 @@ void MstxDomainMarkA(mstxDomainhandle_t handle, const char* message, aclrtStream
185 func(handle, message, stream);189 func(handle, message, stream);
186}190}
187 191 
188-int MstxDomainRangeStartA(mstxDomainhandle_t handle, const char* message, aclrtStream stream, int ptRangeId)192+int MstxDomainRangeStartA(mstxDomainHandle_t handle, const char* message, aclrtStream stream, int ptRangeId)
189{193{
190- using MstxDomainRangeStartAFunc = mstxRangeId (*)(mstxDomainhandle_t, const char*, aclrtStream);194+ using MstxDomainRangeStartAFunc = mstxRangeId (*)(mstxDomainHandle_t, const char*, aclrtStream);
191 static MstxDomainRangeStartAFunc func = nullptr;195 static MstxDomainRangeStartAFunc func = nullptr;
192 static bool noFuncFlag = false;196 static bool noFuncFlag = false;
193 if (noFuncFlag) {197 if (noFuncFlag) {
@@ -207,9 +211,9 @@ int MstxDomainRangeStartA(mstxDomainhandle_t handle, const char* message, aclrtS
207 return 0;211 return 0;
208}212}
209 213 
210-void MstxDomainRangeEnd(mstxDomainhandle_t handle, int ptRangeId)214+void MstxDomainRangeEnd(mstxDomainHandle_t handle, int ptRangeId)
211{215{
212- using MstxDomainRangeEndFunc = void (*)(mstxDomainhandle_t, mstxRangeId);216+ using MstxDomainRangeEndFunc = void (*)(mstxDomainHandle_t, mstxRangeId);
213 static MstxDomainRangeEndFunc func = nullptr;217 static MstxDomainRangeEndFunc func = nullptr;
214 static bool noFuncFlag = false;218 static bool noFuncFlag = false;
215 if (noFuncFlag) {219 if (noFuncFlag) {
@@ -233,5 +237,81 @@ void MstxDomainRangeEnd(mstxDomainhandle_t handle, int ptRangeId)
233 g_rangeIdMap.erase(iter);237 g_rangeIdMap.erase(iter);
234}238}
235 239 
240+mstxMemHeapHandle_t MstxMemHeapRegister(mstxDomainHandle_t domain, mstxMemHeapDesc_t const* desc)
241+{
242+ using MstxMemHeapRegisterFunc = mstxMemHeapHandle_t (*)(mstxDomainHandle_t, mstxMemHeapDesc_t const*);
243+ static MstxMemHeapRegisterFunc func = nullptr;
244+ static bool noFuncFlag = false;
245+ if (noFuncFlag) {
246+ return nullptr;
247+ }
248+ if (func == nullptr) {
249+ func = (MstxMemHeapRegisterFunc)GET_FUNC(mstxMemHeapRegister);
250+ if (func == nullptr) {
251+ ASCEND_LOGW("Failed to get func mstxMemHeapRegister");
252+ noFuncFlag = true;
253+ return nullptr;
254+ }
255+ }
256+ return func(domain, desc);
257+}
258+ 
259+void MstxMemHeapUnregister(mstxDomainHandle_t domain, mstxMemHeapHandle_t heap)
260+{
261+ using MstxMemHeapUnregisterFunc = void (*)(mstxDomainHandle_t, mstxMemHeapHandle_t);
262+ static MstxMemHeapUnregisterFunc func = nullptr;
263+ static bool noFuncFlag = false;
264+ if (noFuncFlag) {
265+ return;
266+ }
267+ if (func == nullptr) {
268+ func = (MstxMemHeapUnregisterFunc)GET_FUNC(mstxMemHeapUnregister);
269+ if (func == nullptr) {
270+ ASCEND_LOGW("Failed to get func mstxMemHeapUnregister");
271+ noFuncFlag = true;
272+ return;
273+ }
274+ }
275+ func(domain, heap);
276+}
277+ 
278+void MstxMemRegionsRegister(mstxDomainHandle_t domain, mstxMemRegionsRegisterBatch_t const* desc)
279+{
280+ using MstxMemRegionsRegisterFunc = void (*)(mstxDomainHandle_t, mstxMemRegionsRegisterBatch_t const*);
281+ static MstxMemRegionsRegisterFunc func = nullptr;
282+ static bool noFuncFlag = false;
283+ if (noFuncFlag) {
284+ return;
285+ }
286+ if (func == nullptr) {
287+ func = (MstxMemRegionsRegisterFunc)GET_FUNC(mstxMemRegionsRegister);
288+ if (func == nullptr) {
289+ ASCEND_LOGW("Failed to get func mstxMemRegionsRegister");
290+ noFuncFlag = true;
291+ return;
292+ }
293+ }
294+ func(domain, desc);
295+}
296+ 
297+void MstxMemRegionsUnregister(mstxDomainHandle_t domain, mstxMemRegionsUnregisterBatch_t const* desc)
298+{
299+ using MstxMemRegionsUnregisterFunc = void (*)(mstxDomainHandle_t, mstxMemRegionsUnregisterBatch_t const*);
300+ static MstxMemRegionsUnregisterFunc func = nullptr;
301+ static bool noFuncFlag = false;
302+ if (noFuncFlag) {
303+ return;
304+ }
305+ if (func == nullptr) {
306+ func = (MstxMemRegionsUnregisterFunc)GET_FUNC(mstxMemRegionsUnregister);
307+ if (func == nullptr) {
308+ ASCEND_LOGW("Failed to get func mstxMemRegionsUnregister");
309+ noFuncFlag = true;
310+ return;
311+ }
312+ }
313+ func(domain, desc);
314+}
315+ 
236}316}
237}317}
@@ -16,15 +16,24 @@ int MstxRangeStartA(const char* message, aclrtStream stream, int ptRangeId);
16 16 
17void MstxRangeEnd(int ptRangeId);17void MstxRangeEnd(int ptRangeId);
18 18 
19-mstxDomainhandle_t MstxDomainCreateA(const char* name);19+mstxDomainHandle_t MstxDomainCreateA(const char* name);
20 20 
21-void MstxDomainDestroy(mstxDomainhandle_t handle);21+void MstxDomainDestroy(mstxDomainHandle_t handle);
22 22 
23-void MstxDomainMarkA(mstxDomainhandle_t handle, const char* message, aclrtStream stream);23+void MstxDomainMarkA(mstxDomainHandle_t handle, const char* message, aclrtStream stream);
24 24 
25-int MstxDomainRangeStartA(mstxDomainhandle_t handle, const char* message, aclrtStream stream, int ptRangeId);25+int MstxDomainRangeStartA(mstxDomainHandle_t handle, const char* message, aclrtStream stream, int ptRangeId);
26+ 
27+void MstxDomainRangeEnd(mstxDomainHandle_t handle, int ptRangeId);
28+ 
29+mstxMemHeapHandle_t MstxMemHeapRegister(mstxDomainHandle_t domain, const mstxMemHeapDesc_t* desc);
30+ 
31+void MstxMemHeapUnregister(mstxDomainHandle_t domain, mstxMemHeapHandle_t heap);
32+ 
33+void MstxMemRegionsRegister(mstxDomainHandle_t domain, const mstxMemRegionsRegisterBatch_t* desc);
34+ 
35+void MstxMemRegionsUnregister(mstxDomainHandle_t domain, const mstxMemRegionsUnregisterBatch_t* desc);
26 36 
27-void MstxDomainRangeEnd(mstxDomainhandle_t handle, int ptRangeId);
28}37}
29}38}
30 39 
@@ -84,17 +84,20 @@ int MstxMgr::getRangeId()
84 return ptRangeId_++;84 return ptRangeId_++;
85}85}
86 86 
87-mstxDomainhandle_t MstxMgr::createDomain(const char* name)87+mstxDomainHandle_t MstxMgr::createDomain(const char* name)
88{88{
89+ if (!isMsleaksEnable() && !isMstxEnable()) {
90+ return nullptr;
91+ }
89 return at_npu::native::MstxDomainCreateA(name);92 return at_npu::native::MstxDomainCreateA(name);
90}93}
91 94 
92-void MstxMgr::destroyDomain(mstxDomainhandle_t domain)95+void MstxMgr::destroyDomain(mstxDomainHandle_t domain)
93{96{
94 at_npu::native::MstxDomainDestroy(domain);97 at_npu::native::MstxDomainDestroy(domain);
95}98}
96 99 
97-void MstxMgr::domainMark(mstxDomainhandle_t domain, const char* message, const aclrtStream stream)100+void MstxMgr::domainMark(mstxDomainHandle_t domain, const char* message, const aclrtStream stream)
98{101{
99 if (!isMstxEnable()) {102 if (!isMstxEnable()) {
100 return;103 return;
@@ -111,7 +114,7 @@ void MstxMgr::domainMark(mstxDomainhandle_t domain, const char* message, const a
111 at_npu::native::OpCommand::RunOpApi("mstx_domain_mark_op", mark_call);114 at_npu::native::OpCommand::RunOpApi("mstx_domain_mark_op", mark_call);
112}115}
113 116 
114-int MstxMgr::domainRangeStart(mstxDomainhandle_t domain, const char* message, const aclrtStream stream)117+int MstxMgr::domainRangeStart(mstxDomainHandle_t domain, const char* message, const aclrtStream stream)
115{118{
116 if (!isMstxEnable()) {119 if (!isMstxEnable()) {
117 return 0;120 return 0;
@@ -133,7 +136,7 @@ int MstxMgr::domainRangeStart(mstxDomainhandle_t domain, const char* message, co
133 return id;136 return id;
134}137}
135 138 
136-void MstxMgr::domainRangeEnd(mstxDomainhandle_t domain, int ptRangeId)139+void MstxMgr::domainRangeEnd(mstxDomainHandle_t domain, int ptRangeId)
137{140{
138 if (!isMstxEnable() || ptRangeId == 0) {141 if (!isMstxEnable() || ptRangeId == 0) {
139 return;142 return;
@@ -158,6 +161,76 @@ void MstxMgr::domainRangeEnd(mstxDomainhandle_t domain, int ptRangeId)
158 at_npu::native::OpCommand::RunOpApi("mstx_domain_range_end_op", range_end_call);161 at_npu::native::OpCommand::RunOpApi("mstx_domain_range_end_op", range_end_call);
159}162}
160 163 
164+mstxMemHeapHandle_t MstxMgr::memHeapRegister(mstxDomainHandle_t domain, mstxMemVirtualRangeDesc_t* desc)
165+{
166+ if (!isMsleaksEnable() || desc==nullptr) {
167+ return nullptr;
168+ }
169+ mstxMemHeapDesc_t heapDesc;
170+ heapDesc.typeSpecificDesc = reinterpret_cast<void const *>(desc);
171+ return at_npu::native::MstxMemHeapRegister(domain, &heapDesc);
172+}
173+ 
174+void MstxMgr::memHeapUnregister(mstxDomainHandle_t domain, void* ptr)
175+{
176+ if (!isMsleaksEnable() || ptr == nullptr) {
177+ return;
178+ }
179+ at_npu::native::MstxMemHeapUnregister(domain, reinterpret_cast<mstxMemHeapHandle_t>(ptr));
180+}
181+ 
182+void MstxMgr::memRegionsRegister(mstxDomainHandle_t domain, mstxMemVirtualRangeDesc_t* desc)
183+{
184+ if (!isMsleaksEnable() || desc == nullptr) {
185+ return;
186+ }
187+ mstxMemRegionsRegisterBatch_t batch;
188+ batch.regionCount = 1;
189+ batch.regionDescArray = reinterpret_cast<const void *>(desc);
190+ at_npu::native::MstxMemRegionsRegister(domain, &batch);
191+}
192+ 
193+void MstxMgr::memRegionsUnregister(mstxDomainHandle_t domain, void* ptr)
194+{
195+ if (!isMsleaksEnable() || ptr == nullptr) {
196+ return;
197+ }
198+ mstxMemRegionsUnregisterBatch_t unregisterBatch;
199+ unregisterBatch.refCount = 1;
200+ mstxMemRegionRef_t regionRef[1] = {};
201+ regionRef[0].refType = MSTX_MEM_REGION_REF_TYPE_POINTER;
202+ regionRef[0].pointer = ptr;
203+ unregisterBatch.refArray = regionRef;
204+ at_npu::native::MstxMemRegionsUnregister(domain, &unregisterBatch);
205+}
206+ 
207+ 
208+bool MstxMgr::isMsleaksEnable()
209+{
210+ static bool isEnable = isMsleaksEnableImpl();
211+ return isEnable;
212+}
213+ 
214+bool MstxMgr::isMsleaksEnableImpl()
215+{
216+ bool ret = false;
217+ const char* envVal = std::getenv("LD_PRELOAD");
218+ if (envVal == nullptr) {
219+ return ret;
220+ }
221+ static const std::string soName = "libascend_hal_hook.so";
222+ std::stringstream ss(envVal);
223+ std::string path;
224+ while (std::getline(ss, path, ':')) {
225+ path = torch_npu::toolkit::profiler::Utils::RealPath(path);
226+ if ((path.size() > soName.size()) && (path.substr(path.size() - soName.size()) == soName)) {
227+ ret = true;
228+ break;
229+ }
230+ }
231+ return ret;
232+}
233+ 
161bool MstxMgr::isProfTxEnable()234bool MstxMgr::isProfTxEnable()
162{235{
163 return ProfilerMgr::GetInstance()->GetNpuTrace().load() && ProfilerMgr::GetInstance()->GetMsprofTx().load();236 return ProfilerMgr::GetInstance()->GetNpuTrace().load() && ProfilerMgr::GetInstance()->GetMsprofTx().load();
@@ -12,6 +12,7 @@ namespace torch_npu {
12namespace profiler {12namespace profiler {
13 13 
14const std::string DOMAIN_COMMUNICATION = "communication";14const std::string DOMAIN_COMMUNICATION = "communication";
15+const std::string DOMAIN_MSLEAKS = "msleaks";
15 16 
16class MstxMgr : public torch_npu::toolkit::profiler::Singleton<MstxMgr> {17class MstxMgr : public torch_npu::toolkit::profiler::Singleton<MstxMgr> {
17friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>;18friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>;
@@ -22,11 +23,15 @@ public:
22 bool isMstxEnable();23 bool isMstxEnable();
23 int getRangeId();24 int getRangeId();
24 25 
25- mstxDomainhandle_t createDomain(const char* name);26+ mstxDomainHandle_t createDomain(const char* name);
26- void destroyDomain(mstxDomainhandle_t domain);27+ void destroyDomain(mstxDomainHandle_t domain);
27- void domainMark(mstxDomainhandle_t domain, const char* message, const aclrtStream stream);28+ void domainMark(mstxDomainHandle_t domain, const char* message, const aclrtStream stream);
28- int domainRangeStart(mstxDomainhandle_t domain, const char* message, const aclrtStream stream);29+ int domainRangeStart(mstxDomainHandle_t domain, const char* message, const aclrtStream stream);
29- void domainRangeEnd(mstxDomainhandle_t domain, int ptRangeId);30+ void domainRangeEnd(mstxDomainHandle_t domain, int ptRangeId);
31+ mstxMemHeapHandle_t memHeapRegister(mstxDomainHandle_t domain, mstxMemVirtualRangeDesc_t* desc);
32+ void memHeapUnregister(mstxDomainHandle_t domain, void* ptr);
33+ void memRegionsRegister(mstxDomainHandle_t domain, mstxMemVirtualRangeDesc_t* desc);
34+ void memRegionsUnregister(mstxDomainHandle_t domain, void* ptr);
30 35 
31private:36private:
32 MstxMgr();37 MstxMgr();
@@ -35,6 +40,8 @@ private:
35 explicit MstxMgr(MstxMgr &&obj) = delete;40 explicit MstxMgr(MstxMgr &&obj) = delete;
36 MstxMgr& operator=(MstxMgr &&obj) = delete;41 MstxMgr& operator=(MstxMgr &&obj) = delete;
37 42 
43+ bool isMsleaksEnable();
44+ bool isMsleaksEnableImpl();
38 bool isProfTxEnable();45 bool isProfTxEnable();
39 bool isMsptiTxEnable();46 bool isMsptiTxEnable();
40 bool isMsptiTxEnableImpl();47 bool isMsptiTxEnableImpl();
@@ -43,6 +50,5 @@ private:
43 std::unordered_set<int> ptRangeIdsWithStream_;50 std::unordered_set<int> ptRangeIdsWithStream_;
44 std::mutex mtx_;51 std::mutex mtx_;
45};52};
46- 
47}53}
48} // namespace torch_npu54} // namespace torch_npu
@@ -129,7 +129,7 @@ inline bool mstxEnable()
129 129 
130struct MstxRange {130struct MstxRange {
131 int rangeId{0};131 int rangeId{0};
132- mstxDomainhandle_t domainHandle{nullptr};132+ mstxDomainHandle_t domainHandle{nullptr};
133 MstxRange(const std::string &message, aclrtStream stream, const std::string &domainName = "default")133 MstxRange(const std::string &message, aclrtStream stream, const std::string &domainName = "default")
134 {134 {
135 if (!mstxEnable()) {135 if (!mstxEnable()) {