[mlir][sparse] use sparse_tensor::StorageSpecifier to store dim/memSizes

Reviewed By: aartbik

Differential Revision: https://reviews.llvm.org/D140130
This commit is contained in:
Peiming Liu
2022-12-15 18:45:07 +00:00
parent 0ebed862d8
commit 988733c600
17 changed files with 678 additions and 1071 deletions

View File

@@ -16,6 +16,7 @@
//===----------------------------------------------------------------------===//
#include "CodegenUtils.h"
#include "SparseTensorStorageLayout.h"
#include "mlir/Dialect/Bufferization/IR/Bufferization.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
@@ -40,38 +41,6 @@ static constexpr const char kInsertFuncNamePrefix[] = "_insert_";
// Helper methods.
//===----------------------------------------------------------------------===//
/// Returns the "tuple" value of the adapted tensor.
static UnrealizedConversionCastOp getTuple(Value tensor) {
return llvm::cast<UnrealizedConversionCastOp>(tensor.getDefiningOp());
}
static SparseTensorDescriptor getDescriptorFromTensorTuple(Value tensor) {
auto tuple = getTuple(tensor);
return SparseTensorDescriptor(tuple.getResultTypes()[0], tuple.getInputs());
}
static MutSparseTensorDescriptor
getMutDescriptorFromTensorTuple(Value tensor, SmallVectorImpl<Value> &fields) {
auto tuple = getTuple(tensor);
fields.assign(tuple.getInputs().begin(), tuple.getInputs().end());
return MutSparseTensorDescriptor(tuple.getResultTypes()[0], fields);
}
/// Packs the given values as a "tuple" value.
static Value genTuple(OpBuilder &builder, Location loc, Type tp,
ValueRange values) {
return builder.create<UnrealizedConversionCastOp>(loc, TypeRange(tp), values)
.getResult(0);
}
static Value genTuple(OpBuilder &builder, Location loc,
SparseTensorDescriptor desc) {
return builder
.create<UnrealizedConversionCastOp>(loc, desc.getTensorType(),
desc.getFields())
.getResult(0);
}
/// Flatten a list of operands that may contain sparse tensors.
static void flattenOperands(ValueRange operands,
SmallVectorImpl<Value> &flattened) {
@@ -146,9 +115,7 @@ static std::optional<Value> sizeFromTensorAtDim(OpBuilder &builder,
// Any other query can consult the dimSizes array at field DimSizesIdx,
// accounting for the reordering applied to the sparse storage.
Value idx = constantIndex(builder, loc, toStoredDim(rtp, dim));
return builder.create<memref::LoadOp>(loc, desc.getDimSizesMemRef(), idx)
.getResult();
return desc.getDimSize(builder, loc, toStoredDim(rtp, dim));
}
// Gets the dimension size at the given stored dimension 'd', either as a
@@ -161,40 +128,24 @@ Value sizeAtStoredDim(OpBuilder &builder, Location loc,
if (!ShapedType::isDynamic(shape[dim]))
return constantIndex(builder, loc, shape[dim]);
return genLoad(builder, loc, desc.getDimSizesMemRef(),
constantIndex(builder, loc, d));
return desc.getDimSize(builder, loc, d);
}
static void createPushback(OpBuilder &builder, Location loc,
MutSparseTensorDescriptor desc, unsigned fidx,
MutSparseTensorDescriptor desc,
SparseTensorFieldKind kind, Optional<unsigned> dim,
Value value, Value repeat = Value()) {
Type etp = desc.getElementType(fidx);
Value field = desc.getField(fidx);
Value newField = builder.create<PushBackOp>(
loc, field.getType(), desc.getMemSizesMemRef(), field,
toType(builder, loc, value, etp), APInt(64, getFieldMemSizesIndex(fidx)),
repeat);
desc.setField(fidx, newField);
}
Type etp = desc.getMemRefElementType(kind, dim);
Value field = desc.getMemRefField(kind, dim);
StorageSpecifierKind specFieldKind = toSpecifierKind(kind);
/// Maps a sparse tensor type to the appropriate compounded buffers.
static std::optional<LogicalResult>
convertSparseTensorType(Type type, SmallVectorImpl<Type> &fields) {
auto enc = getSparseTensorEncoding(type);
if (!enc)
return std::nullopt;
auto pushBackOp = builder.create<PushBackOp>(
loc, desc.getSpecifierField(builder, loc, specFieldKind, dim), field,
toType(builder, loc, value, etp), repeat);
RankedTensorType rType = type.cast<RankedTensorType>();
foreachFieldAndTypeInSparseTensor(
rType,
[&fields](Type fieldType, unsigned fieldIdx,
SparseTensorFieldKind /*fieldKind*/, unsigned /*dim*/,
DimLevelType /*dlt*/) -> bool {
assert(fieldIdx == fields.size());
fields.push_back(fieldType);
return true;
});
return success();
desc.setMemRefField(kind, dim, pushBackOp.getOutBuffer());
desc.setSpecifierField(builder, loc, specFieldKind, dim,
pushBackOp.getNewSize());
}
/// Generates code that allocates a sparse storage scheme for given rank.
@@ -210,8 +161,8 @@ static void allocSchemeForRank(OpBuilder &builder, Location loc,
// the desired "linear + 1" length property at all times.
Type ptrType = getSparseTensorEncoding(rtp).getPointerType();
Value ptrZero = constantZero(builder, loc, ptrType);
createPushback(builder, loc, desc, desc.getPtrMemRefIndex(r), ptrZero,
linear);
createPushback(builder, loc, desc, SparseTensorFieldKind::PtrMemRef, r,
ptrZero, linear);
return;
}
if (isSingletonDim(rtp, r)) {
@@ -226,7 +177,8 @@ static void allocSchemeForRank(OpBuilder &builder, Location loc,
}
// Reached values array so prepare for an insertion.
Value valZero = constantZero(builder, loc, rtp.getElementType());
createPushback(builder, loc, desc, desc.getValMemRefIndex(), valZero, linear);
createPushback(builder, loc, desc, SparseTensorFieldKind::ValMemRef,
std::nullopt, valZero, linear);
}
/// Creates allocation operation.
@@ -257,22 +209,20 @@ static void createAllocFields(OpBuilder &builder, Location loc, Type type,
foreachFieldAndTypeInSparseTensor(
rtp,
[&builder, &fields, loc, heuristic,
[&builder, &fields, rtp, loc, heuristic,
enableInit](Type fType, unsigned fIdx, SparseTensorFieldKind fKind,
unsigned /*dim*/, DimLevelType /*dlt*/) -> bool {
assert(fields.size() == fIdx);
auto memRefTp = fType.cast<MemRefType>();
Value field;
switch (fKind) {
case SparseTensorFieldKind::DimSizes:
case SparseTensorFieldKind::MemSizes:
field = builder.create<memref::AllocOp>(loc, memRefTp);
case SparseTensorFieldKind::StorageSpec:
field = SparseTensorSpecifier::getInitValue(builder, loc, rtp);
break;
case SparseTensorFieldKind::PtrMemRef:
case SparseTensorFieldKind::IdxMemRef:
case SparseTensorFieldKind::ValMemRef:
field =
createAllocation(builder, loc, memRefTp, heuristic, enableInit);
field = createAllocation(builder, loc, fType.cast<MemRefType>(),
heuristic, enableInit);
break;
}
assert(field);
@@ -297,21 +247,18 @@ static void createAllocFields(OpBuilder &builder, Location loc, Type type,
// to all zeros, sets the dimSizes to known values and gives all pointer
// fields an initial zero entry, so that it is easier to maintain the
// "linear + 1" length property.
builder.create<linalg::FillOp>(
loc, constantZero(builder, loc, builder.getIndexType()),
desc.getMemSizesMemRef()); // zero memSizes
Value ptrZero =
constantZero(builder, loc, getSparseTensorEncoding(rtp).getPointerType());
for (unsigned r = 0; r < rank; r++) {
unsigned ro = toOrigDim(rtp, r);
// Fills dim sizes array.
genStore(builder, loc, sizes[ro], desc.getDimSizesMemRef(),
constantIndex(builder, loc, r));
desc.setDimSize(builder, loc, r, sizes[ro]);
// Pushes a leading zero to pointers memref.
if (isCompressedDim(rtp, r))
createPushback(builder, loc, desc, desc.getPtrMemRefIndex(r), ptrZero);
if (isCompressedDim(rtp, r)) {
createPushback(builder, loc, desc, SparseTensorFieldKind::PtrMemRef, r,
ptrZero);
}
}
allocSchemeForRank(builder, loc, desc, /*rank=*/0);
}
@@ -349,10 +296,11 @@ static Value genCompressed(OpBuilder &builder, Location loc,
unsigned ptrIndex = desc.getPtrMemRefIndex(d);
Value one = constantIndex(builder, loc, 1);
Value pp1 = builder.create<arith::AddIOp>(loc, pos, one);
Value plo = genLoad(builder, loc, desc.getField(ptrIndex), pos);
Value phi = genLoad(builder, loc, desc.getField(ptrIndex), pp1);
Value psz = constantIndex(builder, loc, getFieldMemSizesIndex(idxIndex));
Value msz = genLoad(builder, loc, desc.getMemSizesMemRef(), psz);
Value plo = genLoad(builder, loc, desc.getMemRefField(ptrIndex), pos);
Value phi = genLoad(builder, loc, desc.getMemRefField(ptrIndex), pp1);
Value msz = desc.getIdxMemSize(builder, loc, d);
// Value msz = desc.getMemSize(builder, loc, getFieldMemSizesIndex(idxIndex));
Value phim1 = builder.create<arith::SubIOp>(
loc, toType(builder, loc, phi, indexType), one);
// Conditional expression.
@@ -362,14 +310,14 @@ static Value genCompressed(OpBuilder &builder, Location loc,
scf::IfOp ifOp1 = builder.create<scf::IfOp>(loc, types, lt, /*else*/ true);
types.pop_back();
builder.setInsertionPointToStart(&ifOp1.getThenRegion().front());
Value crd = genLoad(builder, loc, desc.getField(idxIndex), phim1);
Value crd = genLoad(builder, loc, desc.getMemRefField(idxIndex), phim1);
Value eq = builder.create<arith::CmpIOp>(loc, arith::CmpIPredicate::eq,
toType(builder, loc, crd, indexType),
indices[d]);
builder.create<scf::YieldOp>(loc, eq);
builder.setInsertionPointToStart(&ifOp1.getElseRegion().front());
if (d > 0)
genStore(builder, loc, msz, desc.getField(ptrIndex), pos);
genStore(builder, loc, msz, desc.getMemRefField(ptrIndex), pos);
builder.create<scf::YieldOp>(loc, constantI1(builder, loc, false));
builder.setInsertionPointAfter(ifOp1);
Value p = ifOp1.getResult(0);
@@ -396,8 +344,9 @@ static Value genCompressed(OpBuilder &builder, Location loc,
// If !present (changes fields, update next).
builder.setInsertionPointToStart(&ifOp2.getElseRegion().front());
Value mszp1 = builder.create<arith::AddIOp>(loc, msz, one);
genStore(builder, loc, mszp1, desc.getField(ptrIndex), pp1);
createPushback(builder, loc, desc, idxIndex, indices[d]);
genStore(builder, loc, mszp1, desc.getMemRefField(ptrIndex), pp1);
createPushback(builder, loc, desc, SparseTensorFieldKind::IdxMemRef, d,
indices[d]);
// Prepare the next dimension "as needed".
if ((d + 1) < rank)
allocSchemeForRank(builder, loc, desc, d + 1);
@@ -459,7 +408,8 @@ static void genInsertBody(OpBuilder &builder, ModuleOp module,
// indices[d].push_back(i[d])
// pos[d] = pos[d-1]
// <insert @ pos[d] at next dimension d + 1>
createPushback(builder, loc, desc, desc.getIdxMemRefIndex(d), indices[d]);
createPushback(builder, loc, desc, SparseTensorFieldKind::IdxMemRef, d,
indices[d]);
} else {
assert(isDenseDim(rtp, d));
// Construct the new position as:
@@ -472,7 +422,8 @@ static void genInsertBody(OpBuilder &builder, ModuleOp module,
}
// Reached the actual value append/insert.
if (!isDenseDim(rtp, rank - 1))
createPushback(builder, loc, desc, desc.getValMemRefIndex(), value);
createPushback(builder, loc, desc, SparseTensorFieldKind::ValMemRef,
std::nullopt, value);
else
genStore(builder, loc, value, desc.getValMemRef(), pos);
builder.create<func::ReturnOp>(loc, fields);
@@ -565,8 +516,7 @@ static void genEndInsert(OpBuilder &builder, Location loc,
if (d > 0) {
Type ptrType = getSparseTensorEncoding(rtp).getPointerType();
Value ptrMemRef = desc.getPtrMemRef(d);
Value mz = constantIndex(builder, loc, desc.getPtrMemSizesIndex(d));
Value hi = genLoad(builder, loc, desc.getMemSizesMemRef(), mz);
Value hi = desc.getPtrMemSize(builder, loc, d);
Value zero = constantIndex(builder, loc, 0);
Value one = constantIndex(builder, loc, 1);
// Vector of only one, but needed by createFor's prototype.
@@ -723,6 +673,7 @@ public:
bool enableInit)
: OpConversionPattern(typeConverter, context),
enableBufferInitialization(enableInit) {}
LogicalResult
matchAndRewrite(bufferization::AllocTensorOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
@@ -761,8 +712,8 @@ public:
// Replace the sparse tensor deallocation with field deallocations.
Location loc = op.getLoc();
auto tuple = getTuple(adaptor.getTensor());
for (auto input : tuple.getInputs())
auto desc = getDescriptorFromTensorTuple(adaptor.getTensor());
for (auto input : desc.getMemRefFields())
// Deallocate every buffer used to store the sparse tensor handler.
rewriter.create<memref::DeallocOp>(loc, input);
@@ -1018,36 +969,13 @@ public:
ConversionPatternRewriter &rewriter) const override {
// Query memSizes for the actually stored values size.
auto desc = getDescriptorFromTensorTuple(adaptor.getTensor());
Value field =
constantIndex(rewriter, op.getLoc(), desc.getValMemSizesIndex());
rewriter.replaceOpWithNewOp<memref::LoadOp>(op, desc.getMemSizesMemRef(),
field);
rewriter.replaceOp(op, desc.getValMemSize(rewriter, op.getLoc()));
return success();
}
};
} // namespace
//===----------------------------------------------------------------------===//
// Sparse tensor type conversion into an actual buffer.
//===----------------------------------------------------------------------===//
mlir::SparseTensorTypeToBufferConverter::SparseTensorTypeToBufferConverter() {
addConversion([](Type type) { return type; });
addConversion(convertSparseTensorType);
// Required by scf.for 1:N type conversion.
addSourceMaterialization([](OpBuilder &builder, RankedTensorType tp,
ValueRange inputs,
Location loc) -> std::optional<Value> {
if (!getSparseTensorEncoding(tp))
// Not a sparse tensor.
return std::nullopt;
// Sparse compiler knows how to cancel out these casts.
return genTuple(builder, loc, tp, inputs);
});
}
//===----------------------------------------------------------------------===//
// Public method for populating conversion rules.
//===----------------------------------------------------------------------===//