[mlir][sparse] implement lowering rules for sparse_tensor.pack operation

Reviewed By: aartbik

Differential Revision: https://reviews.llvm.org/D143230
This commit is contained in:
Peiming Liu
2023-02-03 00:28:12 +00:00
parent 0b8daee028
commit a41672e16a
4 changed files with 209 additions and 8 deletions

View File

@@ -1021,6 +1021,98 @@ public:
}
};
struct SparsePackOpConverter : public OpConversionPattern<PackOp> {
using OpConversionPattern::OpConversionPattern;
LogicalResult
matchAndRewrite(PackOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto rtp = op.getResult().getType().cast<RankedTensorType>();
assert(isUniqueCOOType(rtp));
SmallVector<Value> fields;
Location loc = op.getLoc();
foreachFieldAndTypeInSparseTensor(
rtp,
[&rewriter, &fields, &op, rtp,
loc](Type fType, unsigned fIdx, SparseTensorFieldKind fKind,
unsigned /*dim*/, DimLevelType /*dlt*/) -> bool {
assert(fields.size() == fIdx);
auto enc = getSparseTensorEncoding(rtp);
Value field;
switch (fKind) {
case SparseTensorFieldKind::StorageSpec:
field = SparseTensorSpecifier::getInitValue(rewriter, loc, rtp);
break;
case SparseTensorFieldKind::PtrMemRef: {
// TACO-style COO starts with a PtrBuffer
// By creating a constant value for it, we avoid the complexity of
// memory management.
auto tensorType = RankedTensorType::get({2}, enc.getPointerType());
auto memrefType = MemRefType::get(tensorType.getShape(),
tensorType.getElementType());
auto cstPtr = rewriter.create<arith::ConstantOp>(
loc, tensorType,
DenseElementsAttr::get(
tensorType,
{APInt(64, 0),
APInt(64, op.getData().getType().getShape()[0])}));
field = rewriter.create<bufferization::ToMemrefOp>(loc, memrefType,
cstPtr);
break;
}
case SparseTensorFieldKind::IdxMemRef: {
auto tensorType = op.getIndices().getType();
auto memrefType = MemRefType::get(tensorType.getShape(),
tensorType.getElementType());
auto idxMemRef = rewriter.create<bufferization::ToMemrefOp>(
op->getLoc(), memrefType, op.getIndices());
ReassociationIndices reassociation;
for (int i = 0, e = tensorType.getRank(); i < e; i++)
reassociation.push_back(i);
// Flattened the indices buffer to rank 1.
field = rewriter.create<memref::CollapseShapeOp>(
loc, idxMemRef, ArrayRef<ReassociationIndices>(reassociation));
break;
}
case SparseTensorFieldKind::ValMemRef: {
auto tensorType = op.getData().getType();
auto memrefType = MemRefType::get(tensorType.getShape(),
tensorType.getElementType());
field = rewriter.create<bufferization::ToMemrefOp>(
op->getLoc(), memrefType, op.getData());
break;
}
}
assert(field);
if (fType != field.getType())
field = rewriter.create<memref::CastOp>(loc, fType, field);
fields.push_back(field);
// Returns true to continue the iteration.
return true;
});
MutSparseTensorDescriptor desc(rtp, fields);
auto noe = linalg::createOrFoldDimOp(rewriter, loc, op.getData(), 0);
for (unsigned i = 0, e = rtp.getRank(); i < e; i++) {
int dim = rtp.getShape()[i];
assert(!ShapedType::isDynamic(dim));
desc.setDimSize(rewriter, loc, i, constantIndex(rewriter, loc, dim));
if (i == 0)
desc.setPtrMemSize(rewriter, loc, i, constantIndex(rewriter, loc, 2));
desc.setIdxMemSize(rewriter, loc, i, noe);
}
desc.setValMemSize(rewriter, loc, noe);
rewriter.replaceOp(op, genTuple(rewriter, loc, desc));
return success();
}
};
} // namespace
//===----------------------------------------------------------------------===//
@@ -1032,14 +1124,15 @@ public:
void mlir::populateSparseTensorCodegenPatterns(
TypeConverter &typeConverter, RewritePatternSet &patterns,
bool enableBufferInitialization) {
patterns.add<SparseReturnConverter, SparseCallConverter, SparseDimOpConverter,
SparseCastConverter, SparseTensorDeallocConverter,
SparseTensorLoadConverter, SparseExpandConverter,
SparseCompressConverter, SparseInsertConverter,
SparseToPointersConverter, SparseToIndicesConverter,
SparseToIndicesBufferConverter, SparseToValuesConverter,
SparseConvertConverter, SparseNumberOfEntriesConverter>(
typeConverter, patterns.getContext());
patterns.add<SparsePackOpConverter, SparseReturnConverter,
SparseCallConverter, SparseDimOpConverter, SparseCastConverter,
SparseTensorDeallocConverter, SparseTensorLoadConverter,
SparseExpandConverter, SparseCompressConverter,
SparseInsertConverter, SparseToPointersConverter,
SparseToIndicesConverter, SparseToIndicesBufferConverter,
SparseToValuesConverter, SparseConvertConverter,
SparseNumberOfEntriesConverter>(typeConverter,
patterns.getContext());
patterns.add<SparseTensorAllocConverter>(typeConverter, patterns.getContext(),
enableBufferInitialization);
}