已合并
fix bug:all_gather_v_op.cc存在疑似内存泄露风险 #774
zsj1998创建于 5月8日
fix bug:all_gather_v_op.cc存在疑似内存泄露风险 #774
已合并
zsj1998创建于 5月8日
共 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
Archerls
ArcherlsArcherls6月3日

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

likedislike
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);}};
Archerls
ArcherlsArcherls6月3日

上面是分行写,这里又是一行写的,总感觉很奇怪。这个lambda表达式推荐可以都用一行写

likedislike
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;