149 lines
6.2 KiB
C++
149 lines
6.2 KiB
C++
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
|
#include "mlir/IR/Matchers.h"
|
|
|
|
#include "src/Accelerators/PIM/Conversion/SpatialToPim/ChannelLoweringPatterns.hpp"
|
|
#include "src/Accelerators/PIM/Conversion/SpatialToPim/Common.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 int32_t toPimCoreId(int32_t spatialCoreId) { return spatialCoreId; }
|
|
|
|
static FailureOr<SmallVector<int32_t>> getConstantI32Values(ValueRange values) {
|
|
SmallVector<int32_t> constants;
|
|
constants.reserve(values.size());
|
|
for (Value value : values) {
|
|
APInt constantValue;
|
|
if (!matchPattern(value, m_ConstantInt(&constantValue)))
|
|
return failure();
|
|
constants.push_back(static_cast<int32_t>(constantValue.getSExtValue()));
|
|
}
|
|
return constants;
|
|
}
|
|
|
|
struct ChannelSendLowering : OpRewritePattern<spatial::SpatChannelSendOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatChannelSendOp op, PatternRewriter& rewriter) const override {
|
|
pim::PimSendOp::create(
|
|
rewriter, op.getLoc(), op.getInput(), getTensorSizeInBytesAttr(rewriter, op.getInput()), op.getTargetCoreId());
|
|
rewriter.eraseOp(op);
|
|
return success();
|
|
}
|
|
};
|
|
|
|
struct ChannelReceiveLowering : OpRewritePattern<spatial::SpatChannelReceiveOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatChannelReceiveOp op, PatternRewriter& rewriter) const override {
|
|
if (op->use_empty()) {
|
|
rewriter.eraseOp(op);
|
|
return success();
|
|
}
|
|
auto outputType = cast<ShapedType>(op.getResult().getType());
|
|
Value outputBuffer =
|
|
tensor::EmptyOp::create(rewriter, op.getLoc(), outputType.getShape(), outputType.getElementType()).getResult();
|
|
Value received = pim::PimReceiveOp::create(rewriter,
|
|
op.getLoc(),
|
|
op.getResult().getType(),
|
|
outputBuffer,
|
|
getTensorSizeInBytesAttr(rewriter, op.getResult()),
|
|
op.getSourceCoreId())
|
|
.getOutput();
|
|
rewriter.replaceOp(op, received);
|
|
return success();
|
|
}
|
|
};
|
|
|
|
struct ChannelSendTensorLowering : OpRewritePattern<spatial::SpatChannelSendTensorOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatChannelSendTensorOp op, PatternRewriter& rewriter) const override {
|
|
FailureOr<SmallVector<int32_t>> targetCoreIds = getConstantI32Values(op.getTargetCoreIds());
|
|
if (failed(targetCoreIds))
|
|
return rewriter.notifyMatchFailure(op, "expected constant targetCoreIds");
|
|
for (int32_t& targetCoreId : *targetCoreIds)
|
|
targetCoreId = toPimCoreId(targetCoreId);
|
|
pim::PimSendTensorOp::create(rewriter, op.getLoc(), op.getInput(), rewriter.getDenseI32ArrayAttr(*targetCoreIds));
|
|
rewriter.eraseOp(op);
|
|
return success();
|
|
}
|
|
};
|
|
|
|
struct ChannelReceiveTensorLowering : OpRewritePattern<spatial::SpatChannelReceiveTensorOp> {
|
|
using OpRewritePattern::OpRewritePattern;
|
|
|
|
LogicalResult matchAndRewrite(spatial::SpatChannelReceiveTensorOp op, PatternRewriter& rewriter) const override {
|
|
FailureOr<SmallVector<int32_t>> sourceCoreIds = getConstantI32Values(op.getSourceCoreIds());
|
|
if (failed(sourceCoreIds))
|
|
return rewriter.notifyMatchFailure(op, "expected constant sourceCoreIds");
|
|
for (int32_t& sourceCoreId : *sourceCoreIds)
|
|
sourceCoreId = toPimCoreId(sourceCoreId);
|
|
auto outputType = cast<ShapedType>(op.getOutput().getType());
|
|
Value outputBuffer =
|
|
tensor::EmptyOp::create(rewriter, op.getLoc(), outputType.getShape(), outputType.getElementType()).getResult();
|
|
Value received =
|
|
pim::PimReceiveTensorOp::create(
|
|
rewriter, op.getLoc(), op.getOutput().getType(), outputBuffer, rewriter.getDenseI32ArrayAttr(*sourceCoreIds))
|
|
.getOutput();
|
|
rewriter.replaceOp(op, received);
|
|
return success();
|
|
}
|
|
};
|
|
|
|
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,
|
|
ChannelSendTensorLowering,
|
|
ChannelReceiveTensorLowering,
|
|
ExtractRowsLowering,
|
|
ConcatLowering>(patterns.getContext());
|
|
}
|
|
|
|
} // namespace onnx_mlir
|