diff --git a/src/PIM/Compiler/PimCodeGen.cpp b/src/PIM/Compiler/PimCodeGen.cpp index 9806a59..fa84349 100644 --- a/src/PIM/Compiler/PimCodeGen.cpp +++ b/src/PIM/Compiler/PimCodeGen.cpp @@ -620,10 +620,13 @@ void PimCodeGen::codeGenLmvOp(pim::PimMemCopyOp lmvOp, const StaticValueKnowledg } void PimCodeGen::codeGenReceiveOp(pim::PimReceiveOp receiveOp, const StaticValueKnowledge& knowledge) const { + auto outputOffset = indexOf(receiveOp.getOutputOffset(), knowledge); auto sourceCoreId = indexOf(receiveOp.getSourceCoreId(), knowledge); - assert(succeeded(sourceCoreId) && "pim.receive source core id must be statically resolvable during codegen"); + assert(succeeded(outputOffset) && succeeded(sourceCoreId) + && "pim.receive offset and source core id must be statically resolvable during codegen"); emitCommunicationOp( - pim_binary::Opcode::recv, addressOf(receiveOp.getOutputBuffer(), knowledge), *sourceCoreId, receiveOp.getSize()); + pim_binary::Opcode::recv, addressOf(receiveOp.getOutputBuffer(), knowledge) + *outputOffset, + *sourceCoreId, receiveOp.getSize()); } void PimCodeGen::codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge& knowledge) const { diff --git a/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp b/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp index 1902b83..8b1ad80 100644 --- a/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp +++ b/src/PIM/Conversion/SpatialToPim/CoreLoweringPatterns.cpp @@ -348,7 +348,9 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeOp(spatial::SpatScheduledCom return failure(); Value received = PimReceiveOp::create( - rewriter, receiveOp.getLoc(), outputBuffer.getType(), outputBuffer, *sizeAttr, receiveOp.getSourceCoreId()) + rewriter, receiveOp.getLoc(), outputBuffer.getType(), outputBuffer, + arith::ConstantIndexOp::create(rewriter, receiveOp.getLoc(), 0), + *sizeAttr, receiveOp.getSourceCoreId()) .getOutput(); blockArg->replaceAllUsesWith(received); markOpToRemove(receiveOp); diff --git a/src/PIM/Conversion/SpatialToPim/Patterns/ChannelLowering.cpp b/src/PIM/Conversion/SpatialToPim/Patterns/ChannelLowering.cpp index f46f7de..b1f0e4c 100644 --- a/src/PIM/Conversion/SpatialToPim/Patterns/ChannelLowering.cpp +++ b/src/PIM/Conversion/SpatialToPim/Patterns/ChannelLowering.cpp @@ -85,8 +85,9 @@ struct ChannelReceiveLowering : OpRewritePattern auto sizeAttr = getTensorSizeInBytesAttr(rewriter, op.getOperation(), op.getResult()); if (failed(sizeAttr)) return failure(); + Value zero = arith::ConstantIndexOp::create(rewriter, op.getLoc(), 0); auto receive = pim::PimReceiveOp::create( - rewriter, op.getLoc(), op.getResult().getType(), outputBuffer, *sizeAttr, op.getSourceCoreId()); + rewriter, op.getLoc(), op.getResult().getType(), outputBuffer, zero, *sizeAttr, op.getSourceCoreId()); copyRaptorDebugAttrs(op.getOperation(), receive.getOperation()); Value received = receive.getOutput(); if (!destinationInsert) { @@ -96,7 +97,6 @@ struct ChannelReceiveLowering : OpRewritePattern rewriter.setInsertionPoint(destinationInsert); Value targetOffset = createDestinationByteOffset(rewriter, destinationInsert); - Value zero = arith::ConstantIndexOp::create(rewriter, op.getLoc(), 0); auto copy = pim::PimMemCopyOp::create( rewriter, op.getLoc(), destinationInsert.getDestType(), targetOffset, zero, destinationInsert.getDest(), received, *sizeAttr); diff --git a/src/PIM/Dialect/Pim/Pim.td b/src/PIM/Dialect/Pim/Pim.td index 7db445a..b605385 100644 --- a/src/PIM/Dialect/Pim/Pim.td +++ b/src/PIM/Dialect/Pim/Pim.td @@ -97,6 +97,7 @@ def PimReceiveOp : PimOp<"receive", [DestinationStyleOpInterface]> { let arguments = (ins PimTensor:$outputBuffer, + Index:$outputOffset, I32Attr:$size, Index:$sourceCoreId ); @@ -112,7 +113,8 @@ def PimReceiveOp : PimOp<"receive", [DestinationStyleOpInterface]> { }]; let assemblyFormat = [{ - `(` $outputBuffer `,` $sourceCoreId `)` attr-dict `:` type($outputBuffer) `->` type($output) + `[` $outputOffset `]` `(` $outputBuffer `,` $sourceCoreId `)` attr-dict + `:` type($outputBuffer) `->` type($output) }]; } diff --git a/src/PIM/Dialect/Pim/Transforms/Bufferization/OpBufferizationInterfaces.cpp b/src/PIM/Dialect/Pim/Transforms/Bufferization/OpBufferizationInterfaces.cpp index 0bfd445..69db269 100644 --- a/src/PIM/Dialect/Pim/Transforms/Bufferization/OpBufferizationInterfaces.cpp +++ b/src/PIM/Dialect/Pim/Transforms/Bufferization/OpBufferizationInterfaces.cpp @@ -189,7 +189,8 @@ struct ReceiveOpInterface : DstBufferizableOpInterfaceExternalModelgetLoc(), rewriter); replaceOpWithNewBufferizedOp( - rewriter, op, contiguousOutput.getType(), contiguousOutput, receiveOp.getSizeAttr(), receiveOp.getSourceCoreId()); + rewriter, op, contiguousOutput.getType(), contiguousOutput, receiveOp.getOutputOffset(), + receiveOp.getSizeAttr(), receiveOp.getSourceCoreId()); return success(); } }; diff --git a/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp b/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp index 9f95470..038d428 100644 --- a/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp +++ b/src/PIM/Dialect/Pim/Transforms/Bufferization/PimBufferizationPass.cpp @@ -9,6 +9,7 @@ #include "mlir/IR/Dominance.h" #include "mlir/IR/PatternMatch.h" #include "mlir/Interfaces/DestinationStyleOpInterface.h" +#include "mlir/Interfaces/SideEffectInterfaces.h" #include "mlir/Pass/Pass.h" #include "mlir/Transforms/GreedyPatternRewriteDriver.h" @@ -177,12 +178,22 @@ static void forwardSingleConsumerContiguousInputCopies(func::FuncOp funcOp) { if (hasOtherUse || !consumerUse) continue; - Value output = getForwardedInputConsumerOutput(*consumerUse); Value source = copy.getSource(); - if (!output || !isDeviceLocalPimAddress(source) + if (!isDeviceLocalPimAddress(source) || (failed(resolveContiguousAddress(source)) && failed(compileContiguousAddressExpr(source)))) continue; + if (isa(consumerUse->getOwner())) { + consumerUse->set(source); + copy.erase(); + targetAlloc.erase(); + continue; + } + + Value output = getForwardedInputConsumerOutput(*consumerUse); + if (!output) + continue; + FailureOr sourceBase = getPimAddressBase(source); FailureOr outputBase = getPimAddressBase(output); if (failed(sourceBase) || failed(outputBase) || *sourceBase == *outputBase) @@ -251,6 +262,46 @@ static void forwardSingleConsumerPimOutputCopies(func::FuncOp funcOp) { } } +static void forwardSingleConsumerReceiveCopies(func::FuncOp funcOp) { + SmallVector copies; + funcOp.walk([&](PimMemCopyOp copy) { copies.push_back(copy); }); + + for (PimMemCopyOp copy : copies) { + auto receive = copy.getSource().getDefiningOp(); + auto outputAlloc = receive + ? receive.getOutputBuffer().getDefiningOp() + : memref::AllocOp(); + auto sourceOffset = getConstantIntValue(copy.getSourceOffset()); + if (!receive || !outputAlloc || !receive.getOutput().hasOneUse() + || !outputAlloc.getResult().hasOneUse() + || !sourceOffset || *sourceOffset != 0 + || copy.getSize() != receive.getSize() + || receive->getBlock() != copy->getBlock()) + continue; + + bool canSinkReceive = true; + for (Operation* between = receive->getNextNode(); between != copy; + between = between->getNextNode()) { + if (!isMemoryEffectFree(between)) { + canSinkReceive = false; + break; + } + } + if (!canSinkReceive) + continue; + + OpBuilder builder(copy); + auto forwarded = PimReceiveOp::create( + builder, receive.getLoc(), copy.getOutput().getType(), copy.getTarget(), + copy.getTargetOffset(), receive.getSizeAttr(), receive.getSourceCoreId()); + forwarded->setAttrs(receive->getAttrs()); + copy.getOutput().replaceAllUsesWith(forwarded.getOutput()); + copy.erase(); + receive.erase(); + outputAlloc.erase(); + } +} + enum class ExpectedPimCopyDirection { HostToDevice, DeviceToHost, DeviceToDevice }; static LogicalResult verifyPimCopyEndpoints(Operation* copy, @@ -406,6 +457,7 @@ void PimBufferizationPass::runOnOperation() { return; } + forwardSingleConsumerReceiveCopies(funcOp); forwardSingleConsumerContiguousInputCopies(funcOp); forwardSingleConsumerPimOutputCopies(funcOp); diff --git a/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp b/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp index 7a3cae7..d032c0c 100644 --- a/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp +++ b/src/PIM/Dialect/Pim/Transforms/Verification/VerificationPass.cpp @@ -990,6 +990,15 @@ private: hasFailure = true; } } + if (auto receiveOp = dyn_cast(op); + receiveOp + && failed(resolveIndexValue(receiveOp.getOutputOffset(), knowledge))) { + diagnostics.report(&op, [](Operation* illegalOp) { + illegalOp->emitOpError( + "output offset must be statically evaluable for PIM codegen"); + }); + hasFailure = true; + } return success(!hasFailure); }); }