#include "mlir/Dialect/Arith/IR/Arith.h" #include "mlir/Dialect/Tensor/IR/Tensor.h" #include "src/Accelerators/PIM/Conversion/SpatialToPim/Common.hpp" #include "src/Accelerators/PIM/Conversion/SpatialToPim/Patterns.hpp" #include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp" #include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp" #include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp" using namespace mlir; namespace onnx_mlir { namespace { static void copyRaptorDebugAttrs(Operation* source, Operation* target) { for (NamedAttribute attr : source->getAttrs()) { StringRef name = attr.getName().strref(); if (name.starts_with("raptor.")) target->setAttr(attr.getName(), attr.getValue()); } } static Value createDestinationByteOffset(PatternRewriter& rewriter, tensor::InsertSliceOp insert) { auto destinationType = cast(insert.getDestType()); SmallVector strides = computeRowMajorStrides(destinationType.getShape()); int64_t elementBytes = getElementTypeSizeInBytes(destinationType.getElementType()); Value total = arith::ConstantIndexOp::create(rewriter, insert.getLoc(), 0); for (auto [dimension, offset] : llvm::enumerate(insert.getMixedOffsets())) { int64_t scale = strides[dimension] * elementBytes; Value component; if (auto attribute = dyn_cast(offset)) { component = arith::ConstantIndexOp::create( rewriter, insert.getLoc(), cast(attribute).getInt() * scale); } else { component = cast(offset); if (scale != 1) component = arith::MulIOp::create( rewriter, insert.getLoc(), component, arith::ConstantIndexOp::create(rewriter, insert.getLoc(), scale)); } total = arith::AddIOp::create(rewriter, insert.getLoc(), total, component); } return total; } struct ChannelSendLowering : OpRewritePattern { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(spatial::SpatChannelSendOp op, PatternRewriter& rewriter) const override { auto sizeAttr = getTensorSizeInBytesAttr(rewriter, op.getOperation(), op.getInput()); if (failed(sizeAttr)) return failure(); auto send = pim::PimSendOp::create(rewriter, op.getLoc(), op.getInput(), *sizeAttr, op.getTargetCoreId()); copyRaptorDebugAttrs(op.getOperation(), send.getOperation()); rewriter.eraseOp(op); return success(); } }; struct HostStoreSyncLowering : OpRewritePattern { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(spatial::SpatHostStoreSyncOp op, PatternRewriter& rewriter) const override { auto sizeAttr = getTensorSizeInBytesAttr(rewriter, op.getOperation(), op.getInput()); auto hostBuffer = getPipelineHostBuffer(rewriter, op); if (failed(sizeAttr) || failed(hostBuffer)) return failure(); Value zero = arith::ConstantIndexOp::create(rewriter, op.getLoc(), 0); pim::PimMemCopyDevToHostOp::create( rewriter, op.getLoc(), hostBuffer->getType(), op.getHostOffset(), zero, *hostBuffer, op.getInput(), *sizeAttr); auto sync = pim::PimSyncOp::create( rewriter, op.getLoc(), op.getTargetCoreId(), op.getEventRegister()); copyRaptorDebugAttrs(op.getOperation(), sync.getOperation()); rewriter.eraseOp(op); return success(); } }; template static LogicalResult lowerReceive( ReceiveOp op, PatternRewriter& rewriter, CreateReceive createReceive) { if (op->use_empty()) { rewriter.eraseOp(op); return success(); } auto outputType = cast(op.getResult().getType()); tensor::InsertSliceOp destinationInsert; if (op->hasOneUse()) { auto insert = dyn_cast(*op->getUsers().begin()); auto destinationType = insert ? dyn_cast(insert.getDestType()) : RankedTensorType(); if (insert && insert.getSource() == op.getOutput() && insert.getSourceType() == outputType && insert->getBlock() == op->getBlock() && destinationType && destinationType.hasStaticShape() && isContiguousSubviewWithDynamicOffsets( destinationType.getShape(), insert.getMixedOffsets(), insert.getStaticSizes(), insert.getStaticStrides())) destinationInsert = insert; } Value outputBuffer = tensor::EmptyOp::create(rewriter, op.getLoc(), outputType.getShape(), outputType.getElementType()).getResult(); auto sizeAttr = getTensorSizeInBytesAttr(rewriter, op.getOperation(), op.getResult()); if (failed(sizeAttr)) return failure(); Value zero = arith::ConstantIndexOp::create(rewriter, op.getLoc(), 0); auto received = createReceive(outputBuffer, zero, *sizeAttr); if (failed(received)) return failure(); if (!destinationInsert) { rewriter.replaceOp(op, *received); return success(); } rewriter.setInsertionPoint(destinationInsert); Value targetOffset = createDestinationByteOffset(rewriter, destinationInsert); auto copy = pim::PimMemCopyOp::create( rewriter, op.getLoc(), destinationInsert.getDestType(), targetOffset, zero, destinationInsert.getDest(), *received, *sizeAttr); rewriter.replaceOp(destinationInsert, copy.getOutput()); rewriter.eraseOp(op); return success(); } struct ChannelReceiveLowering : OpRewritePattern { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(spatial::SpatChannelReceiveOp op, PatternRewriter& rewriter) const override { return lowerReceive(op, rewriter, [&](Value outputBuffer, Value zero, IntegerAttr sizeAttr) -> FailureOr { auto receive = pim::PimReceiveOp::create( rewriter, op.getLoc(), op.getResult().getType(), outputBuffer, zero, sizeAttr, op.getSourceCoreId()); copyRaptorDebugAttrs(op.getOperation(), receive.getOperation()); return receive.getOutput(); }); } }; struct HostWaitLoadLowering : OpRewritePattern { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(spatial::SpatHostWaitLoadOp op, PatternRewriter& rewriter) const override { return lowerReceive(op, rewriter, [&](Value outputBuffer, Value zero, IntegerAttr sizeAttr) -> FailureOr { auto hostBuffer = getPipelineHostBuffer(rewriter, op); if (failed(hostBuffer)) return failure(); auto wait = pim::PimWaitOp::create( rewriter, op.getLoc(), op.getEventRegister(), rewriter.getI32IntegerAttr(1)); copyRaptorDebugAttrs(op.getOperation(), wait.getOperation()); return pim::PimMemCopyHostToDevOp::create( rewriter, op.getLoc(), outputBuffer.getType(), zero, op.getHostOffset(), outputBuffer, *hostBuffer, sizeAttr).getOutput(); }); } }; struct ExtractRowsLowering : OpRewritePattern { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(spatial::SpatExtractRowsOp op, PatternRewriter& rewriter) const override { auto inputType = cast(op.getInput().getType()); SmallVector replacements; replacements.reserve(op.getNumResults()); for (auto [rowIndex, output] : llvm::enumerate(op.getOutputs())) { auto outputType = cast(output.getType()); SmallVector offsets = { rewriter.getIndexAttr(static_cast(rowIndex) * outputType.getDimSize(0)), rewriter.getIndexAttr(0)}; SmallVector sizes = {rewriter.getIndexAttr(outputType.getDimSize(0)), rewriter.getIndexAttr(inputType.getDimSize(1))}; SmallVector strides = {rewriter.getIndexAttr(1), rewriter.getIndexAttr(1)}; replacements.push_back( tensor::ExtractSliceOp::create(rewriter, op.getLoc(), outputType, op.getInput(), offsets, sizes, strides) .getResult()); } rewriter.replaceOp(op, replacements); return success(); } }; struct ConcatLowering : OpRewritePattern { using OpRewritePattern::OpRewritePattern; LogicalResult matchAndRewrite(spatial::SpatConcatOp op, PatternRewriter& rewriter) const override { auto outputType = cast(op.getOutput().getType()); Value outputBuffer = tensor::EmptyOp::create(rewriter, op.getLoc(), outputType.getShape(), outputType.getElementType()).getResult(); Value concatenated = pim::PimConcatOp::create( rewriter, op.getLoc(), op.getOutput().getType(), op.getAxisAttr(), op.getInputs(), outputBuffer) .getOutput(); rewriter.replaceOp(op, concatenated); return success(); } }; } // namespace void populateChannelLoweringPatterns(RewritePatternSet& patterns) { patterns.add(patterns.getContext()); } } // namespace onnx_mlir