@@ -198,7 +198,7 @@ void EpDispatchCombineHandle::InitializeShmemBuf() {
198198 config.WeightBytes () + config.SrcTokenIdBytes () + blockwiseScaleBytes);
199199 }
200200
201- if (config.kernelType == KernelType::IntraNode) {
201+ if (config.kernelType == KernelType::IntraNode || config. kernelType == KernelType::IntraNodeLL ) {
202202 auto & bufs = shmemTokBufs.emplace <ShmemBufsIntraNode>();
203203 bufs.combineInp = ShmemMallocAndReturnMemObjPtr (maxStagingSize, hipDeviceMallocUncached);
204204 bufs.dispatchOut = ShmemMallocAndReturnMemObjPtr (dispatchOutSize, hipDeviceMallocUncached);
@@ -271,7 +271,7 @@ void EpDispatchCombineHandle::InitializeShmemBuf() {
271271}
272272
273273void EpDispatchCombineHandle::FinalizeShmemBuf () {
274- if (config.kernelType == KernelType::IntraNode) {
274+ if (config.kernelType == KernelType::IntraNode || config. kernelType == KernelType::IntraNodeLL ) {
275275 auto & bufs = std::get<ShmemBufsIntraNode>(shmemTokBufs);
276276 ShmemFree (bufs.dispatchOut ->localPtr );
277277 ShmemFree (bufs.combineInp ->localPtr );
@@ -463,7 +463,8 @@ EpDispatchCombineArgsRaw GetEpDispatchCombineArgsRaw(const EpDispatchCombineHand
463463 args.scalesBuf = handle.scalesBuf ;
464464 args.destPeTokenCounter = handle.destPeTokenCounter ;
465465 args.localPeTokenCounter = handle.localPeTokenCounter ;
466- if (handle.config .kernelType == KernelType::IntraNode) {
466+ if (handle.config .kernelType == KernelType::IntraNode ||
467+ handle.config .kernelType == KernelType::IntraNodeLL) {
467468 args.intraNodeTokBufs = std::get<ShmemBufsIntraNode>(handle.shmemTokBufs );
468469 } else if (handle.config .kernelType == KernelType::InterNodeV1 ||
469470 handle.config .kernelType == KernelType::InterNodeV1LL) {
0 commit comments