已合并
Add mstx dotting to obtain PTA memory pool data #18242
czy6创建于 2025年2月23日
Add mstx dotting to obtain PTA memory pool data #18242
已合并
从refs/pull/18242/head合入到master
共 8 个文件变更+293-36
| @@ -13,7 +13,61 @@ typedef uint64_t mstxRangeId; | |||
| 13 | 13 | ||
| 14 | struct mstxDomainRegistration_st; | 14 | struct mstxDomainRegistration_st; |
| 15 | typedef struct mstxDomainRegistration_st mstxDomainRegistration_t; | 15 | typedef 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 | ||
| 18 | ACL_FUNC_VISIBILITY void mstxMarkA(const char* message, aclrtStream stream); | 72 | ACL_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 | ||
| 22 | ACL_FUNC_VISIBILITY void mstxRangeEnd(mstxRangeId id); | 76 | ACL_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, |
| 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 | 97 | ||
| 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 | 1263 | ||
| 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 | 1325 | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | 55 | ||
| 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 | 93 | ||
| 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 | 138 | ||
| 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) | |||
| 24 | LOAD_FUNCTION(mstxDomainMarkA) | 24 | LOAD_FUNCTION(mstxDomainMarkA) |
| 25 | LOAD_FUNCTION(mstxDomainRangeStartA) | 25 | LOAD_FUNCTION(mstxDomainRangeStartA) |
| 26 | LOAD_FUNCTION(mstxDomainRangeEnd) | 26 | LOAD_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 invalid | 33 | // 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 | ||
| 17 | void MstxRangeEnd(int ptRangeId); | 17 | void 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 | + | ||
| 161 | bool MstxMgr::isProfTxEnable() | 234 | bool 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 { | |||
| 12 | namespace profiler { | 12 | namespace profiler { |
| 13 | 13 | ||
| 14 | const std::string DOMAIN_COMMUNICATION = "communication"; | 14 | const std::string DOMAIN_COMMUNICATION = "communication"; |
| 15 | +const std::string DOMAIN_MSLEAKS = "msleaks"; | ||
| 15 | 16 | ||
| 16 | class MstxMgr : public torch_npu::toolkit::profiler::Singleton<MstxMgr> { | 17 | class MstxMgr : public torch_npu::toolkit::profiler::Singleton<MstxMgr> { |
| 17 | friend class torch_npu::toolkit::profiler::Singleton<MstxMgr>; | 18 | friend 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 | ||
| 31 | private: | 36 | private: |
| 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_npu | 54 | } // namespace torch_npu |
| @@ -129,7 +129,7 @@ inline bool mstxEnable() | |||
| 129 | 129 | ||
| 130 | struct MstxRange { | 130 | struct 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()) { |
同版本修改,不存在兼容性问题