已合并
fix bug:all_gather_v_op.cc存在疑似内存泄露风险 #774
zsj1998创建于 5月8日
fix bug:all_gather_v_op.cc存在疑似内存泄露风险 #774
已合并
共 1 个文件变更+13-10
| @@ -167,7 +167,14 @@ HcclResult AllGatherVOutPlace(void *sendBuf, void *recvBuf, uint64_t sendCount,c | |||
| 167 | HCCL_ERROR("malloc OpParam failed!"); | 167 | HCCL_ERROR("malloc OpParam failed!"); |
| 168 | return HCCL_E_INTERNAL; | 168 | return HCCL_E_INTERNAL; |
| 169 | } | 169 | } |
| 170 | - OpParam* paramPtr = new (paramMem) OpParam(); | 170 | + OpParam* tmpParamPtr = new (paramMem) OpParam(); |
| 171 | + auto deleter = [](OpParam* p) { | ||
| 172 | + if (p) { | ||
| 173 | + p->~OpParam(); | ||
| 174 | + free(p); | ||
| 175 | + } | ||
| 176 | + }; | ||
| 177 | + std::unique_ptr<OpParam, decltype(deleter)> paramPtr(tmpParamPtr, deleter); | ||
| 171 | OpParam& param = *paramPtr; | 178 | OpParam& param = *paramPtr; |
| 172 | CHK_RET(HcclGetCommName(comm, param.commName)); | 179 | CHK_RET(HcclGetCommName(comm, param.commName)); |
| 173 | param.opMode = OpMode::OPBASE; | 180 | param.opMode = OpMode::OPBASE; |
| @@ -219,8 +226,6 @@ HcclResult AllGatherVOutPlace(void *sendBuf, void *recvBuf, uint64_t sendCount,c | |||
| 219 | return HcclAllGatherVInner(sendBuf, sendCount, recvBuf, recvCounts, recvDispls, dataType, comm, stream); | 226 | return HcclAllGatherVInner(sendBuf, sendCount, recvBuf, recvCounts, recvDispls, dataType, comm, stream); |
| 220 | } | 227 | } |
| 221 | CHK_RET(HcclExecOp(comm, param, topoInfo, algName)); | 228 | CHK_RET(HcclExecOp(comm, param, topoInfo, algName)); |
| 222 | - paramPtr->~OpParam(); | ||
| 223 | - free(paramMem); | ||
| 224 | HCCL_INFO("Execute AllGatherVOutPlace success."); | 229 | HCCL_INFO("Execute AllGatherVOutPlace success."); |
| 225 | return HCCL_SUCCESS; | 230 | return HCCL_SUCCESS; |
| 226 | } | 231 | } |
| @@ -250,16 +255,12 @@ HcclResult AllGatherVOutPlaceGraphMode(void *sendBuf, void *recvBuf, uint64_t se | |||
| 250 | HCCL_INFO("Start to execute AllGatherVOutPlaceGraphMode"); | 255 | HCCL_INFO("Start to execute AllGatherVOutPlaceGraphMode"); |
| 251 | u32 userRankSize; | 256 | u32 userRankSize; |
| 252 | CHK_RET(HcclGetRankSize(comm, &userRankSize)); | 257 | CHK_RET(HcclGetRankSize(comm, &userRankSize)); |
| 253 | - | ||
| 254 | u32 perDataSize = DATATYPE_SIZE_TABLE[dataType]; | 258 | u32 perDataSize = DATATYPE_SIZE_TABLE[dataType]; |
| 255 | u64 inputSize = sendCount * perDataSize; // all gather v 每个rank上一份数据 | 259 | u64 inputSize = sendCount * perDataSize; // all gather v 每个rank上一份数据 |
| 256 | u64 outputSize = 0; | 260 | u64 outputSize = 0; |
| 257 | const u64 *u64RecvCount = reinterpret_cast<const u64 *>(recvCounts); | 261 | const u64 *u64RecvCount = reinterpret_cast<const u64 *>(recvCounts); |
| 258 | const u64 *u64RecvDispls = reinterpret_cast<const u64 *>(recvDispls); | 262 | const u64 *u64RecvDispls = reinterpret_cast<const u64 *>(recvDispls); |
| 259 | - for (u64 i = 0; i < userRankSize; i++) { | 263 | + for (u64 i = 0; i < userRankSize; i++) {outputSize = (outputSize > (u64RecvDispls[i] + u64RecvCount[i]) * perDataSize) ? outputSize : (u64RecvDispls[i] + u64RecvCount[i]) * perDataSize;}// 结果为最大的displs加recvcount |
| 260 | - outputSize = (outputSize > (u64RecvDispls[i] + u64RecvCount[i]) * perDataSize) ? outputSize : (u64RecvDispls[i] + u64RecvCount[i]) * perDataSize; | ||
| 261 | - }// 结果为最大的displs加recvcount | ||
| 262 | - // 申请OpParam参数结构体内存 | ||
| 263 | u64 varMemSize = (userRankSize + userRankSize) * sizeof(u64); | 264 | u64 varMemSize = (userRankSize + userRankSize) * sizeof(u64); |
| 264 | void* paramMem = malloc(sizeof(OpParam) + varMemSize); | 265 | void* paramMem = malloc(sizeof(OpParam) + varMemSize); |
| 265 | if (!paramMem) { | 266 | if (!paramMem) { |
| @@ -267,8 +268,10 @@ HcclResult AllGatherVOutPlaceGraphMode(void *sendBuf, void *recvBuf, uint64_t se | |||
| 267 | HCCL_ERROR("malloc OpParam failed!"); | 268 | HCCL_ERROR("malloc OpParam failed!"); |
| 268 | return HCCL_E_INTERNAL; | 269 | return HCCL_E_INTERNAL; |
| 269 | } | 270 | } |
| 270 | - OpParam* paramPtr = new (paramMem) OpParam(); | 271 | + OpParam* tmpParamPtr = new (paramMem) OpParam(); |
| 271 | - OpParam& param = *paramPtr; | 272 | + auto deleter = [](OpParam* p) {if (p) {p->~OpParam(); free(p);}}; |
上面是分行写,这里又是一行写的,总感觉很奇怪。这个lambda表达式推荐可以都用一行写 ![]() ![]() | |||
| 273 | + std::unique_ptr<OpParam, decltype(deleter)> paramPtr(tmpParamPtr, deleter); | ||
| 274 | + OpParam& param = *paramPtr; | ||
| 272 | CHK_RET(HcclGetCommName(comm, param.commName)); | 275 | CHK_RET(HcclGetCommName(comm, param.commName)); |
| 273 | 276 | ||
| 274 | DevType deviceType = DevType::DEV_TYPE_COUNT; | 277 | DevType deviceType = DevType::DEV_TYPE_COUNT; |


这里代码不对吧,写成一行?