207 lines
8.9 KiB
C++
207 lines
8.9 KiB
C++
#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<RankedTensorType>(insert.getDestType());
|
|
SmallVector<int64_t> 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<Attribute>(offset)) {
|
|
component = arith::ConstantIndexOp::create(
|
|
rewriter, insert.getLoc(), cast<IntegerAttr>(attribute).getInt() * scale);
|
|
} else {
|
|
component = cast<Value>(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<spatial::SpatChannelSendOp> {
|
|
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<spatial::SpatHostStoreSyncOp> {
|
|
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 <typename ReceiveOp, typename CreateReceive>
|
|
static LogicalResult lowerReceive(
|
|
ReceiveOp op, PatternRewriter& rewriter, CreateReceive createReceive) {
|
|
if (op->use_empty()) {
|
|
rewriter.eraseOp(op);
|
|
return success();
|
|
}
|
|
auto outputType = cast<RankedTensorType>(op.getResult().getType());
|
|
tensor::InsertSliceOp destinationInsert;
|
|
if (op->hasOneUse()) {
|
|
auto insert = dyn_cast<tensor::InsertSliceOp>(*op->getUsers().begin());
|
|
auto destinationType = insert
|
|
? dyn_cast<RankedTensorType>(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<spatial::SpatChannelReceiveOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatChannelReceiveOp op, PatternRewriter& rewriter) const override {
|
|
return lowerReceive(op, rewriter, [&](Value outputBuffer, Value zero, IntegerAttr sizeAttr) -> FailureOr<Value> {
|
|
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<spatial::SpatHostWaitLoadOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatHostWaitLoadOp op, PatternRewriter& rewriter) const override {
|
|
return lowerReceive(op, rewriter, [&](Value outputBuffer, Value zero, IntegerAttr sizeAttr) -> FailureOr<Value> {
|
|
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<spatial::SpatExtractRowsOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatExtractRowsOp op, PatternRewriter& rewriter) const override {
|
|
auto inputType = cast<RankedTensorType>(op.getInput().getType());
|
|
SmallVector<Value> replacements;
|
|
replacements.reserve(op.getNumResults());
|
|
for (auto [rowIndex, output] : llvm::enumerate(op.getOutputs())) {
|
|
auto outputType = cast<RankedTensorType>(output.getType());
|
|
SmallVector<OpFoldResult> offsets = {
|
|
rewriter.getIndexAttr(static_cast<int64_t>(rowIndex) * outputType.getDimSize(0)), rewriter.getIndexAttr(0)};
|
|
SmallVector<OpFoldResult> sizes = {rewriter.getIndexAttr(outputType.getDimSize(0)),
|
|
rewriter.getIndexAttr(inputType.getDimSize(1))};
|
|
SmallVector<OpFoldResult> 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<spatial::SpatConcatOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatConcatOp op, PatternRewriter& rewriter) const override {
|
|
auto outputType = cast<ShapedType>(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<ChannelSendLowering, ChannelReceiveLowering,
|
|
HostStoreSyncLowering, HostWaitLoadLowering,
|
|
ExtractRowsLowering, ConcatLowering>(patterns.getContext());
|
|
}
|
|
|
|
} // namespace onnx_mlir
|