[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:
@@ -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.
|
||||
//===----------------------------------------------------------------------===//
|
||||
|
||||
Reference in New Issue
Block a user