[mlir][sparse] unifies sparse_tensor.sort_coo/sort into one operation. (#66722)
The use cases of the two operations are largely overlapped, let's simplify it and only use one of them.
This commit is contained in:
@@ -1353,35 +1353,15 @@ LogicalResult SelectOp::verify() {
|
||||
return success();
|
||||
}
|
||||
|
||||
LogicalResult SortOp::verify() {
|
||||
if (getXs().empty())
|
||||
return emitError("need at least one xs buffer.");
|
||||
|
||||
std::optional<int64_t> n = getConstantIntValue(getN());
|
||||
|
||||
Type xtp = getMemRefType(getXs().front()).getElementType();
|
||||
auto checkTypes = [&](ValueRange operands,
|
||||
bool checkEleType = true) -> LogicalResult {
|
||||
for (Value opnd : operands) {
|
||||
auto mtp = getMemRefType(opnd);
|
||||
const DynSize sh = mtp.getShape()[0];
|
||||
// We can't check the size of dynamic dimension at compile-time, but all
|
||||
// xs and ys should have a dimension not less than n at runtime.
|
||||
if (n && !ShapedType::isDynamic(sh) && sh < n.value())
|
||||
return emitError(llvm::formatv("xs and ys need to have a dimension >= n"
|
||||
": {0} < {1}",
|
||||
sh, n.value()));
|
||||
|
||||
if (checkEleType && xtp != mtp.getElementType())
|
||||
return emitError("mismatch xs element types");
|
||||
}
|
||||
return success();
|
||||
};
|
||||
RETURN_FAILURE_IF_FAILED(checkTypes(getXs()))
|
||||
return n ? checkTypes(getYs(), false) : success();
|
||||
}
|
||||
|
||||
LogicalResult SortCooOp::verify() {
|
||||
AffineMap xPerm = getPermMap();
|
||||
uint64_t nx = xPerm.getNumDims();
|
||||
if (nx < 1)
|
||||
emitError(llvm::formatv("Expected rank(perm_map) > 1, got {0}", nx));
|
||||
|
||||
if (!xPerm.isPermutation())
|
||||
emitError(llvm::formatv("Expected a permutation map, got {0}", xPerm));
|
||||
|
||||
std::optional<int64_t> cn = getConstantIntValue(getN());
|
||||
// We can't check the size of the buffers when n or buffer dimensions aren't
|
||||
// compile-time constants.
|
||||
@@ -1389,12 +1369,6 @@ LogicalResult SortCooOp::verify() {
|
||||
return success();
|
||||
|
||||
uint64_t n = cn.value();
|
||||
uint64_t nx = 1;
|
||||
if (auto nxAttr = getNxAttr()) {
|
||||
nx = nxAttr.getInt();
|
||||
if (nx < 1)
|
||||
emitError(llvm::formatv("Expected nx > 1, got {0}", nx));
|
||||
}
|
||||
uint64_t ny = 0;
|
||||
if (auto nyAttr = getNyAttr()) {
|
||||
ny = nyAttr.getInt();
|
||||
@@ -1409,7 +1383,8 @@ LogicalResult SortCooOp::verify() {
|
||||
emitError(llvm::formatv("{0} got {1} < {2}", message, sh, minSize));
|
||||
};
|
||||
|
||||
checkDim(getXy(), n * (nx + ny), "Expected dimension(xy) >= n * (nx + ny)");
|
||||
checkDim(getXy(), n * (nx + ny),
|
||||
"Expected dimension(xy) >= n * (rank(perm_map) + ny)");
|
||||
|
||||
for (Value opnd : getYs()) {
|
||||
checkDim(opnd, n, "Expected dimension(y) >= n");
|
||||
|
||||
Reference in New Issue
Block a user