[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:
@@ -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);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user