[Mlir-commits] [mlir] [mlir][linalg] Add maps-based inferConvolutionDims overload (PR #203323)
Pooja Hemashekar
llvmlistbot at llvm.org
Fri Jun 12 09:31:06 PDT 2026
https://github.com/phemashekar updated https://github.com/llvm/llvm-project/pull/203323
>From 1e19a2317f420408d9f6c2b8687ba135a0bf3267 Mon Sep 17 00:00:00 2001
From: Pooja Hemashekar <hemashekar at roofline.ai>
Date: Wed, 10 Jun 2026 22:31:55 +0000
Subject: [PATCH 1/5] [mlir][linalg] Add maps-based inferConvolutionDims
overload
This adds an overload of `inferConvolutionDims` that takes only indexing
maps (input, filter, output), mirroring the maps-based
`inferContractionDims` overload added in #76081, along with the
corresponding C API and python bindings.
---
mlir/include/mlir-c/Dialect/Linalg.h | 4 +
.../mlir/Dialect/Linalg/IR/LinalgInterfaces.h | 12 ++
mlir/lib/Bindings/Python/DialectLinalg.cpp | 32 +++++
mlir/lib/CAPI/Dialect/Linalg.cpp | 42 ++++++
.../Dialect/Linalg/IR/LinalgInterfaces.cpp | 68 ++++++---
mlir/test/python/dialects/linalg/utils.py | 51 +++++++
.../Linalg/InferConvolutionDimsTest.cpp | 135 ++++++++++++++++++
7 files changed, 324 insertions(+), 20 deletions(-)
diff --git a/mlir/include/mlir-c/Dialect/Linalg.h b/mlir/include/mlir-c/Dialect/Linalg.h
index 003b0cde39652..b75b9e9490f00 100644
--- a/mlir/include/mlir-c/Dialect/Linalg.h
+++ b/mlir/include/mlir-c/Dialect/Linalg.h
@@ -55,6 +55,10 @@ typedef struct MlirLinalgConvolutionDimensions {
MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
mlirLinalgInferConvolutionDimensions(MlirOperation op);
+MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
+mlirLinalgInferConvolutionDimensionsFromMaps(const MlirAffineMap *indexingMaps,
+ size_t numMaps);
+
MLIR_CAPI_EXPORTED MlirAttribute
mlirLinalgGetIndexingMapsAttribute(MlirOperation op);
diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgInterfaces.h b/mlir/include/mlir/Dialect/Linalg/IR/LinalgInterfaces.h
index 3c7ebd8277dbd..bed70816d7f1e 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgInterfaces.h
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgInterfaces.h
@@ -114,6 +114,18 @@ struct ConvolutionDimensions {
/// Returns a failure if `output_image` (and implicitly `filter_loop`) is empty.
FailureOr<ConvolutionDimensions> inferConvolutionDims(LinalgOp linalgOp);
+/// Maps-based overload of `inferConvolutionDims`. The `indexingMaps` are
+/// expected in operand order: input, filter, output. The iterator types are
+/// inferred from the output map: dimensions that appear in the output are
+/// parallel, all others are reduction. Since there is no operation to carry
+/// native `strides`/`dilations` attributes, the strides and dilations are
+/// derived from the convolution access pattern in the input indexing map.
+/// Returns a failure if there are not exactly 3 maps, the output map is not a
+/// projected permutation, or `output_image` (and implicitly `filter_loop`) is
+/// empty.
+FailureOr<ConvolutionDimensions>
+inferConvolutionDims(ArrayRef<AffineMap> indexingMaps);
+
/// Checks whether `linalgOp` conforms to ConvolutionOpInterface.
/// By default, we require the `linalgOp` to have non-empty convolved dims
/// (implicitly non-empty `output_image` and `filter_loop`).
diff --git a/mlir/lib/Bindings/Python/DialectLinalg.cpp b/mlir/lib/Bindings/Python/DialectLinalg.cpp
index 299961e100786..f53aaae58e442 100644
--- a/mlir/lib/Bindings/Python/DialectLinalg.cpp
+++ b/mlir/lib/Bindings/Python/DialectLinalg.cpp
@@ -178,6 +178,38 @@ static void populateDialectLinalgSubmodule(nb::module_ m) {
m.def("infer_convolution_dimensions", &InferConvolutionDimensions,
"Infers convolution dimensions", nb::arg("op"));
+ m.def(
+ "infer_convolution_dimensions_from_maps",
+ [](std::vector<PyAffineMap> indexingMaps)
+ -> std::optional<PyLinalgConvolutionDimensions> {
+ if (indexingMaps.empty())
+ return std::nullopt;
+
+ std::vector<MlirAffineMap> indexingMaps_(indexingMaps.size());
+ std::copy(indexingMaps.begin(), indexingMaps.end(),
+ indexingMaps_.begin());
+ MlirLinalgConvolutionDimensions dims =
+ mlirLinalgInferConvolutionDimensionsFromMaps(indexingMaps_.data(),
+ indexingMaps_.size());
+
+ // Detect "empty" result from invalid input or failed inference.
+ if (mlirAttributeIsNull(dims.batch) &&
+ mlirAttributeIsNull(dims.outputImage) &&
+ mlirAttributeIsNull(dims.outputChannel) &&
+ mlirAttributeIsNull(dims.filterLoop) &&
+ mlirAttributeIsNull(dims.inputChannel) &&
+ mlirAttributeIsNull(dims.depth) &&
+ mlirAttributeIsNull(dims.strides) &&
+ mlirAttributeIsNull(dims.dilations)) {
+ return std::nullopt;
+ }
+ return dims;
+ },
+ "Infers convolution dimensions (batch/output_image/output_channel/"
+ "filter_loop/input_channel/depth/strides/dilations) from a list of "
+ "affine maps.",
+ nb::arg("indexing_maps"));
+
m.def(
"get_indexing_maps",
[](PyOperationBase &op) -> std::optional<PyArrayAttribute> {
diff --git a/mlir/lib/CAPI/Dialect/Linalg.cpp b/mlir/lib/CAPI/Dialect/Linalg.cpp
index 75c811aed6cc5..24acb983d264f 100644
--- a/mlir/lib/CAPI/Dialect/Linalg.cpp
+++ b/mlir/lib/CAPI/Dialect/Linalg.cpp
@@ -152,6 +152,48 @@ mlirLinalgInferConvolutionDimensions(MlirOperation op) {
return result;
}
+MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
+mlirLinalgInferConvolutionDimensionsFromMaps(const MlirAffineMap *indexingMaps,
+ size_t numMaps) {
+ MlirLinalgConvolutionDimensions result{};
+ if (!indexingMaps || numMaps == 0)
+ return result;
+
+ SmallVector<AffineMap, 3> maps;
+ maps.reserve(numMaps);
+ for (size_t i = 0; i < numMaps; ++i)
+ maps.push_back(unwrap(indexingMaps[i]));
+
+ FailureOr<linalg::ConvolutionDimensions> maybeDims =
+ linalg::inferConvolutionDims(maps);
+ if (failed(maybeDims))
+ return result;
+
+ const linalg::ConvolutionDimensions &dims = *maybeDims;
+ MLIRContext *ctx = maps[0].getContext();
+
+ auto toI32Attr =
+ [&ctx](const SmallVector<unsigned, 2> &vals) -> MlirAttribute {
+ return wrap(DenseI32ArrayAttr::get(ctx, llvm::to_vector_of<int32_t>(vals)));
+ };
+
+ auto toI64Attr =
+ [&ctx](const SmallVector<int64_t, 2> &vals) -> MlirAttribute {
+ return wrap(DenseI64ArrayAttr::get(ctx, vals));
+ };
+
+ result.batch = toI32Attr(dims.batch);
+ result.outputImage = toI32Attr(dims.outputImage);
+ result.outputChannel = toI32Attr(dims.outputChannel);
+ result.filterLoop = toI32Attr(dims.filterLoop);
+ result.inputChannel = toI32Attr(dims.inputChannel);
+ result.depth = toI32Attr(dims.depth);
+ result.strides = toI64Attr(dims.strides);
+ result.dilations = toI64Attr(dims.dilations);
+
+ return result;
+}
+
MLIR_CAPI_EXPORTED MlirAttribute
mlirLinalgGetIndexingMapsAttribute(MlirOperation op) {
auto linalgOp = llvm::dyn_cast<mlir::linalg::LinalgOp>(unwrap(op));
diff --git a/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp b/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp
index 238bddcb3b2bd..f3984177dba2f 100644
--- a/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp
+++ b/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp
@@ -761,18 +761,21 @@ getConstantsFromExprList(const SmallVector<AffineExpr, 2> &exprs) {
/// - `strides[i]` corresponds to `outputImage[i]`.
/// - `dilations[i]` corresponds to `filterLoop[i]`.
/// - Other dimension sets (batch, outputChannel, etc.) are sorted by index.
-static FailureOr<ConvolutionDimensions>
-inferConvolutionDimsImpl(LinalgOp linalgOp,
- ConvAccessExprWalker &inputExprWalker,
- bool allowEmptyConvolvedDims) {
- auto filterMap =
- linalgOp.getMatchingIndexingMap(linalgOp.getDpsInputOperand(1));
- auto outputMap =
- linalgOp.getMatchingIndexingMap(linalgOp.getDpsInitOperand(0));
- llvm::SmallDenseSet<int64_t> filterDims = findPermutationsIndexingOperand(
- filterMap, linalgOp.getIteratorTypesArray(), par);
- llvm::SmallDenseSet<int64_t> outputDims = findPermutationsIndexingOperand(
- outputMap, linalgOp.getIteratorTypesArray(), par);
+///
+/// `nativeStrides` and `nativeDilations`, when non-null, are the op-carried
+/// `strides`/`dilations` attributes and take precedence over the values derived
+/// from the convolution access pattern. They are null for the maps-based
+/// overload.
+static FailureOr<ConvolutionDimensions> inferConvolutionDimsImpl(
+ ArrayRef<AffineMap> indexingMaps, ArrayRef<utils::IteratorType> iterators,
+ ConvAccessExprWalker &inputExprWalker, bool allowEmptyConvolvedDims,
+ DenseIntElementsAttr nativeStrides, DenseIntElementsAttr nativeDilations) {
+ AffineMap filterMap = indexingMaps[1];
+ AffineMap outputMap = indexingMaps.back();
+ llvm::SmallDenseSet<int64_t> filterDims =
+ findPermutationsIndexingOperand(filterMap, iterators, par);
+ llvm::SmallDenseSet<int64_t> outputDims =
+ findPermutationsIndexingOperand(outputMap, iterators, par);
// unConvolvedDims & outputDims - filterDims are the batch iterators.
llvm::SmallDenseSet<int64_t> batch = inputExprWalker.unConvolvedDims;
@@ -794,8 +797,7 @@ inferConvolutionDimsImpl(LinalgOp linalgOp,
llvm::set_intersect(depth, inputExprWalker.unConvolvedDims);
llvm::SmallDenseSet<int64_t> filterReducedDims =
- findPermutationsIndexingOperand(filterMap,
- linalgOp.getIteratorTypesArray(), red);
+ findPermutationsIndexingOperand(filterMap, iterators, red);
// convolvedDims & filterReducedDims are the filter loop iterators.
llvm::SmallDenseSet<int64_t> fl = inputExprWalker.convolvedDims;
@@ -832,7 +834,6 @@ inferConvolutionDimsImpl(LinalgOp linalgOp,
dimensions.filterLoop.push_back(inputExprWalker.convolvedDimMapping[oiDim]);
// Use the op carried strides/dilations attribute if present.
- auto nativeStrides = linalgOp->getAttrOfType<DenseIntElementsAttr>("strides");
if (!nativeStrides) {
SmallVector<AffineExpr, 2> strideExprs;
for (unsigned oiDim : dimensions.outputImage)
@@ -841,8 +842,6 @@ inferConvolutionDimsImpl(LinalgOp linalgOp,
} else {
dimensions.strides = llvm::to_vector<2>(nativeStrides.getValues<int64_t>());
}
- auto nativeDilations =
- linalgOp->getAttrOfType<DenseIntElementsAttr>("dilations");
if (!nativeDilations) {
SmallVector<AffineExpr, 2> dilationExprs;
for (unsigned flDim : dimensions.filterLoop)
@@ -896,8 +895,35 @@ mlir::linalg::inferConvolutionDims(LinalgOp linalgOp) {
(void)inputExprWalker.visit(expr);
inputExprWalker.clearMultiUseDims(indexingMaps[0]);
- return inferConvolutionDimsImpl(linalgOp, inputExprWalker,
- /*allowEmptyConvolvedDims=*/false);
+ return inferConvolutionDimsImpl(
+ indexingMaps, linalgOp.getIteratorTypesArray(), inputExprWalker,
+ /*allowEmptyConvolvedDims=*/false,
+ linalgOp->getAttrOfType<DenseIntElementsAttr>("strides"),
+ linalgOp->getAttrOfType<DenseIntElementsAttr>("dilations"));
+}
+
+FailureOr<ConvolutionDimensions>
+mlir::linalg::inferConvolutionDims(ArrayRef<AffineMap> indexingMaps) {
+ if (indexingMaps.size() != 3)
+ return failure();
+
+ // Infer iterator types from the output map.
+ FailureOr<SmallVector<utils::IteratorType>> iterators =
+ inferIteratorsFromOutMap(indexingMaps[2]);
+ if (failed(iterators))
+ return failure();
+
+ // Check the input indexing map has the right form.
+ ConvAccessExprWalker inputExprWalker;
+ for (AffineExpr expr : indexingMaps[0].getResults())
+ (void)inputExprWalker.visit(expr);
+ inputExprWalker.clearMultiUseDims(indexingMaps[0]);
+
+ return inferConvolutionDimsImpl(indexingMaps, iterators.value(),
+ inputExprWalker,
+ /*allowEmptyConvolvedDims=*/false,
+ /*nativeStrides=*/nullptr,
+ /*nativeDilations=*/nullptr);
}
namespace mlir::linalg::detail {
@@ -1040,7 +1066,9 @@ mlir::linalg::detail::isConvolutionInterfaceImpl(
if (dimensions) {
FailureOr<ConvolutionDimensions> res = inferConvolutionDimsImpl(
- linalgOp, inputExprWalker, allowEmptyConvolvedDims);
+ indexingMaps, iteratorTypes, inputExprWalker, allowEmptyConvolvedDims,
+ linalgOp->getAttrOfType<DenseIntElementsAttr>("strides"),
+ linalgOp->getAttrOfType<DenseIntElementsAttr>("dilations"));
assert(succeeded(res) && "unexpected failure to infer convolution dims");
*dimensions = *res;
}
diff --git a/mlir/test/python/dialects/linalg/utils.py b/mlir/test/python/dialects/linalg/utils.py
index 8ab53b4e28743..e91b2408e1382 100644
--- a/mlir/test/python/dialects/linalg/utils.py
+++ b/mlir/test/python/dialects/linalg/utils.py
@@ -248,3 +248,54 @@ def test_infer_contraction_dimensions_from_maps():
assert len(elementwise_dims.n) == 0
assert len(elementwise_dims.k) == 0
assert list(elementwise_dims.batch) == [0, 1]
+
+
+ at run
+def test_infer_convolution_dimensions_from_maps():
+ with Context(), Location.unknown():
+ module = Module.create()
+ with InsertionPoint(module.body):
+ # === NHWC/HWCF convolution ===
+ # Loop order: (n=d0, oh=d1, ow=d2, f=d3, kh=d4, kw=d5, c=d6).
+ d0 = AffineDimExpr.get(0)
+ d1 = AffineDimExpr.get(1)
+ d2 = AffineDimExpr.get(2)
+ d3 = AffineDimExpr.get(3)
+ d4 = AffineDimExpr.get(4)
+ d5 = AffineDimExpr.get(5)
+ d6 = AffineDimExpr.get(6)
+
+ input_map = AffineMap.get(7, 0, [d0, d1 + d4, d2 + d5, d6])
+ filter_map = AffineMap.get(7, 0, [d4, d5, d6, d3])
+ output_map = AffineMap.get(7, 0, [d0, d1, d2, d3])
+
+ dims = linalg.infer_convolution_dimensions_from_maps(
+ [input_map, filter_map, output_map]
+ )
+ assert dims is not None
+ assert list(dims.batch) == [0]
+ assert list(dims.output_image) == [1, 2]
+ assert list(dims.output_channel) == [3]
+ assert list(dims.filter_loop) == [4, 5]
+ assert list(dims.input_channel) == [6]
+ assert list(dims.depth) == []
+ assert list(dims.strides) == [1, 1]
+ assert list(dims.dilations) == [1, 1]
+
+ # === Invalid input (wrong number of maps) ===
+ invalid_dims = linalg.infer_convolution_dimensions_from_maps(
+ [input_map, filter_map]
+ )
+ assert invalid_dims is None
+
+ # === Non-convolution (matmul-like) returns None ===
+ dim_m = AffineDimExpr.get(0)
+ dim_n = AffineDimExpr.get(1)
+ dim_k = AffineDimExpr.get(2)
+ a_map = AffineMap.get(3, 0, [dim_m, dim_k])
+ b_map = AffineMap.get(3, 0, [dim_k, dim_n])
+ c_map = AffineMap.get(3, 0, [dim_m, dim_n])
+ non_conv = linalg.infer_convolution_dimensions_from_maps(
+ [a_map, b_map, c_map]
+ )
+ assert non_conv is None
diff --git a/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp b/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp
index 7f495a4859064..b5a3747660a3f 100644
--- a/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp
+++ b/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp
@@ -176,4 +176,139 @@ TEST_F(InferConvolutionDimsTest, Conv2DPairing) {
<< "outputImage[1]=1 should pair with filterLoop[1]=2 (ow <-> kw)";
}
+/// Asserts that two ConvolutionDimensions are equal across every populated
+/// field.
+static void expectConvDimsEq(const ConvolutionDimensions &lhs,
+ const ConvolutionDimensions &rhs) {
+ EXPECT_EQ(lhs.batch, rhs.batch);
+ EXPECT_EQ(lhs.outputImage, rhs.outputImage);
+ EXPECT_EQ(lhs.outputChannel, rhs.outputChannel);
+ EXPECT_EQ(lhs.filterLoop, rhs.filterLoop);
+ EXPECT_EQ(lhs.inputChannel, rhs.inputChannel);
+ EXPECT_EQ(lhs.depth, rhs.depth);
+ EXPECT_EQ(lhs.strides, rhs.strides);
+ EXPECT_EQ(lhs.dilations, rhs.dilations);
+}
+
+/// Verify that inferring convolution dimensions from indexing maps produces
+/// same result as inferring thrm directly from the convolution op.
+TEST_F(InferConvolutionDimsTest, MapsOverloadMatchesOpOverload) {
+ OpBuilder builder(ctx.get());
+ OwningOpRef<ModuleOp> module = ModuleOp::create(builder.getUnknownLoc());
+ builder.setInsertionPointToStart(module->getBody());
+ Location loc = builder.getUnknownLoc();
+ Type f32 = builder.getF32Type();
+ auto empty = [&](ArrayRef<int64_t> shape) -> Value {
+ return tensor::EmptyOp::create(builder, loc, shape, f32);
+ };
+ auto ones = builder.getI64TensorAttr({1, 1});
+
+ SmallVector<Operation *> convs;
+ {
+ Value out = empty({1, 4, 4, 4});
+ convs.push_back(linalg::Conv2DNhwcHwcfOp::create(
+ builder, loc, out.getType(),
+ ValueRange{empty({1, 10, 5, 3}), empty({2, 2, 3, 4})},
+ ValueRange{out},
+ /*strides=*/builder.getI64TensorAttr({2, 1}),
+ /*dilations=*/builder.getI64TensorAttr({3, 1}))
+ .getOperation());
+ }
+ {
+ Value out = empty({1, 4, 4, 3});
+ convs.push_back(linalg::DepthwiseConv2DNhwcHwcOp::create(
+ builder, loc, out.getType(),
+ ValueRange{empty({1, 5, 5, 3}), empty({2, 2, 3})},
+ ValueRange{out}, ones, ones)
+ .getOperation());
+ }
+
+ for (Operation *op : convs) {
+ LinalgOp linalgOp = cast<LinalgOp>(op);
+ FailureOr<ConvolutionDimensions> fromOp = inferConvolutionDims(linalgOp);
+ ASSERT_TRUE(succeeded(fromOp))
+ << "op overload failed for " << op->getName().getStringRef().str();
+ FailureOr<ConvolutionDimensions> fromMaps =
+ inferConvolutionDims(linalgOp.getIndexingMapsArray());
+ ASSERT_TRUE(succeeded(fromMaps))
+ << "maps overload failed for " << op->getName().getStringRef().str();
+ expectConvDimsEq(*fromOp, *fromMaps);
+ }
+
+ // Ensure the depthwise conv populates the depth field.
+ FailureOr<ConvolutionDimensions> depthwiseDims =
+ inferConvolutionDims(cast<LinalgOp>(convs.back()).getIndexingMapsArray());
+ ASSERT_TRUE(succeeded(depthwiseDims));
+ EXPECT_FALSE(depthwiseDims->depth.empty());
+}
+
+/// The maps overload must infer the correct dimensions for a strided and
+/// dilated NHWC/HWCF convolution.
+TEST_F(InferConvolutionDimsTest, InferNhwcConvDimensions) {
+ MLIRContext *c = ctx.get();
+ AffineExpr d0, d1, d2, d3, d4, d5, d6;
+ bindDims(c, d0, d1, d2, d3, d4, d5, d6);
+ // Loop order: (n=d0, oh=d1, ow=d2, f=d3, kh=d4, kw=d5, c=d6),
+ // strides {2,1}, dilations {3,1}.
+ SmallVector<AffineMap> maps = {
+ AffineMap::get(7, 0, {d0, d1 * 2 + d4 * 3, d2 + d5, d6}, c),
+ AffineMap::get(7, 0, {d4, d5, d6, d3}, c),
+ AffineMap::get(7, 0, {d0, d1, d2, d3}, c)};
+ FailureOr<ConvolutionDimensions> dims = inferConvolutionDims(maps);
+ ASSERT_TRUE(succeeded(dims));
+ ConvolutionDimensions expected{/*batch=*/{0},
+ /*outputImage=*/{1, 2},
+ /*outputChannel=*/{3},
+ /*filterLoop=*/{4, 5},
+ /*inputChannel=*/{6},
+ /*depth=*/{},
+ /*strides=*/{2, 1},
+ /*dilations=*/{3, 1}};
+ expectConvDimsEq(expected, *dims);
+}
+
+/// Maps-only check that filterLoop pairing is reconstructed from the input map:
+/// the swapped layout (d0 + d3, d1 + d2) must yield filterLoop [3, 2].
+TEST_F(InferConvolutionDimsTest, MapsOverloadConv2DPairing) {
+ MLIRContext *c = ctx.get();
+ AffineExpr d0, d1, d2, d3;
+ bindDims(c, d0, d1, d2, d3);
+ SmallVector<AffineMap> maps = {AffineMap::get(4, 0, {d0 + d3, d1 + d2}, c),
+ AffineMap::get(4, 0, {d2, d3}, c),
+ AffineMap::get(4, 0, {d0, d1}, c)};
+ FailureOr<ConvolutionDimensions> dims = inferConvolutionDims(maps);
+ ASSERT_TRUE(succeeded(dims));
+ EXPECT_EQ(dims->outputImage, (SmallVector<unsigned, 2>{0, 1}));
+ EXPECT_EQ(dims->filterLoop, (SmallVector<unsigned, 2>{3, 2}));
+}
+
+/// Inputs that cannot describe a convolution must fail.
+TEST_F(InferConvolutionDimsTest, MapsOverloadRejectsInvalidInput) {
+ MLIRContext *c = ctx.get();
+ AffineExpr d0, d1, d2, d3;
+ bindDims(c, d0, d1, d2, d3);
+
+ // Wrong number of maps (must be exactly 3: input, filter, output).
+ {
+ auto m = AffineMap::get(2, 0, {d0, d1}, c);
+ EXPECT_TRUE(failed(inferConvolutionDims(SmallVector<AffineMap>{m, m})));
+ }
+
+ // Output map is not a projected permutation: iterator types are unknowable.
+ {
+ SmallVector<AffineMap> maps = {AffineMap::get(4, 0, {d0 + d2, d1 + d3}, c),
+ AffineMap::get(4, 0, {d2, d3}, c),
+ AffineMap::get(4, 0, {d0 + d1}, c)};
+ EXPECT_TRUE(failed(inferConvolutionDims(maps)));
+ }
+
+ // No convolved dimension (matmul-like): output_image is empty.
+ {
+ SmallVector<AffineMap> maps = {AffineMap::get(3, 0, {d0, d2}, c),
+ AffineMap::get(3, 0, {d2, d1}, c),
+ AffineMap::get(3, 0, {d0, d1}, c)};
+ EXPECT_TRUE(failed(inferConvolutionDims(maps)));
+ }
+}
+
} // namespace
>From c36f0fd6e467839e2a978dbf380f9a762ba22e7b Mon Sep 17 00:00:00 2001
From: Pooja Hemashekar <hemashekar at roofline.ai>
Date: Fri, 12 Jun 2026 08:49:40 +0000
Subject: [PATCH 2/5] address comments
---
mlir/lib/CAPI/Dialect/Linalg.cpp | 73 +++++++++-----------------------
1 file changed, 21 insertions(+), 52 deletions(-)
diff --git a/mlir/lib/CAPI/Dialect/Linalg.cpp b/mlir/lib/CAPI/Dialect/Linalg.cpp
index 24acb983d264f..3b50df0b218af 100644
--- a/mlir/lib/CAPI/Dialect/Linalg.cpp
+++ b/mlir/lib/CAPI/Dialect/Linalg.cpp
@@ -115,49 +115,40 @@ MLIR_CAPI_EXPORTED bool mlirLinalgIsAConvolutionOp(MlirOperation op) {
return linalg::isaConvolutionOpInterface(linalgOp);
}
+static MlirLinalgConvolutionDimensions
+toConvolutionDimensions(MLIRContext *ctx,
+ const linalg::ConvolutionDimensions &dims) {
+ auto toI32Attr = [ctx](ArrayRef<unsigned> vals) -> MlirAttribute {
+ return wrap(DenseI32ArrayAttr::get(ctx, llvm::to_vector_of<int32_t>(vals)));
+ };
+ auto toI64Attr = [ctx](ArrayRef<int64_t> vals) -> MlirAttribute {
+ return wrap(DenseI64ArrayAttr::get(ctx, vals));
+ };
+ return {toI32Attr(dims.batch), toI32Attr(dims.outputImage),
+ toI32Attr(dims.outputChannel), toI32Attr(dims.filterLoop),
+ toI32Attr(dims.inputChannel), toI32Attr(dims.depth),
+ toI64Attr(dims.strides), toI64Attr(dims.dilations)};
+}
+
MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
mlirLinalgInferConvolutionDimensions(MlirOperation op) {
- MlirLinalgConvolutionDimensions result{};
auto linalgOp = llvm::dyn_cast<mlir::linalg::LinalgOp>(unwrap(op));
if (!linalgOp)
- return result;
+ return MlirLinalgConvolutionDimensions{};
FailureOr<linalg::ConvolutionDimensions> maybeDims =
linalg::inferConvolutionDims(linalgOp);
if (failed(maybeDims))
- return result;
-
- const linalg::ConvolutionDimensions &dims = *maybeDims;
- MLIRContext *ctx = linalgOp.getContext();
-
- auto toI32Attr =
- [&ctx](const SmallVector<unsigned, 2> &vals) -> MlirAttribute {
- return wrap(DenseI32ArrayAttr::get(ctx, llvm::to_vector_of<int32_t>(vals)));
- };
-
- auto toI64Attr =
- [&ctx](const SmallVector<int64_t, 2> &vals) -> MlirAttribute {
- return wrap(DenseI64ArrayAttr::get(ctx, vals));
- };
-
- result.batch = toI32Attr(dims.batch);
- result.outputImage = toI32Attr(dims.outputImage);
- result.outputChannel = toI32Attr(dims.outputChannel);
- result.filterLoop = toI32Attr(dims.filterLoop);
- result.inputChannel = toI32Attr(dims.inputChannel);
- result.depth = toI32Attr(dims.depth);
- result.strides = toI64Attr(dims.strides);
- result.dilations = toI64Attr(dims.dilations);
+ return MlirLinalgConvolutionDimensions{};
- return result;
+ return toConvolutionDimensions(linalgOp.getContext(), *maybeDims);
}
MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
mlirLinalgInferConvolutionDimensionsFromMaps(const MlirAffineMap *indexingMaps,
size_t numMaps) {
- MlirLinalgConvolutionDimensions result{};
if (!indexingMaps || numMaps == 0)
- return result;
+ return MlirLinalgConvolutionDimensions{};
SmallVector<AffineMap, 3> maps;
maps.reserve(numMaps);
@@ -167,31 +158,9 @@ mlirLinalgInferConvolutionDimensionsFromMaps(const MlirAffineMap *indexingMaps,
FailureOr<linalg::ConvolutionDimensions> maybeDims =
linalg::inferConvolutionDims(maps);
if (failed(maybeDims))
- return result;
-
- const linalg::ConvolutionDimensions &dims = *maybeDims;
- MLIRContext *ctx = maps[0].getContext();
+ return MlirLinalgConvolutionDimensions{};
- auto toI32Attr =
- [&ctx](const SmallVector<unsigned, 2> &vals) -> MlirAttribute {
- return wrap(DenseI32ArrayAttr::get(ctx, llvm::to_vector_of<int32_t>(vals)));
- };
-
- auto toI64Attr =
- [&ctx](const SmallVector<int64_t, 2> &vals) -> MlirAttribute {
- return wrap(DenseI64ArrayAttr::get(ctx, vals));
- };
-
- result.batch = toI32Attr(dims.batch);
- result.outputImage = toI32Attr(dims.outputImage);
- result.outputChannel = toI32Attr(dims.outputChannel);
- result.filterLoop = toI32Attr(dims.filterLoop);
- result.inputChannel = toI32Attr(dims.inputChannel);
- result.depth = toI32Attr(dims.depth);
- result.strides = toI64Attr(dims.strides);
- result.dilations = toI64Attr(dims.dilations);
-
- return result;
+ return toConvolutionDimensions(maps[0].getContext(), *maybeDims);
}
MLIR_CAPI_EXPORTED MlirAttribute
>From d684ad208a4874ea9bdb27555391b34832941f64 Mon Sep 17 00:00:00 2001
From: Pooja Hemashekar <hemashekar at roofline.ai>
Date: Fri, 12 Jun 2026 10:18:29 +0000
Subject: [PATCH 3/5] style nits
---
mlir/lib/CAPI/Dialect/Linalg.cpp | 9 ++++-----
1 file changed, 4 insertions(+), 5 deletions(-)
diff --git a/mlir/lib/CAPI/Dialect/Linalg.cpp b/mlir/lib/CAPI/Dialect/Linalg.cpp
index 3b50df0b218af..dd5955ed3c6e8 100644
--- a/mlir/lib/CAPI/Dialect/Linalg.cpp
+++ b/mlir/lib/CAPI/Dialect/Linalg.cpp
@@ -134,12 +134,12 @@ MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
mlirLinalgInferConvolutionDimensions(MlirOperation op) {
auto linalgOp = llvm::dyn_cast<mlir::linalg::LinalgOp>(unwrap(op));
if (!linalgOp)
- return MlirLinalgConvolutionDimensions{};
+ return {};
FailureOr<linalg::ConvolutionDimensions> maybeDims =
linalg::inferConvolutionDims(linalgOp);
if (failed(maybeDims))
- return MlirLinalgConvolutionDimensions{};
+ return {};
return toConvolutionDimensions(linalgOp.getContext(), *maybeDims);
}
@@ -148,17 +148,16 @@ MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
mlirLinalgInferConvolutionDimensionsFromMaps(const MlirAffineMap *indexingMaps,
size_t numMaps) {
if (!indexingMaps || numMaps == 0)
- return MlirLinalgConvolutionDimensions{};
+ return {};
SmallVector<AffineMap, 3> maps;
- maps.reserve(numMaps);
for (size_t i = 0; i < numMaps; ++i)
maps.push_back(unwrap(indexingMaps[i]));
FailureOr<linalg::ConvolutionDimensions> maybeDims =
linalg::inferConvolutionDims(maps);
if (failed(maybeDims))
- return MlirLinalgConvolutionDimensions{};
+ return {};
return toConvolutionDimensions(maps[0].getContext(), *maybeDims);
}
>From 60d74121e78e8c45c548761a309aa921de632bcd Mon Sep 17 00:00:00 2001
From: Pooja Hemashekar <hemashekar at roofline.ai>
Date: Fri, 12 Jun 2026 11:11:33 +0000
Subject: [PATCH 4/5] address comments
---
mlir/lib/CAPI/Dialect/Linalg.cpp | 4 +++-
1 file changed, 3 insertions(+), 1 deletion(-)
diff --git a/mlir/lib/CAPI/Dialect/Linalg.cpp b/mlir/lib/CAPI/Dialect/Linalg.cpp
index dd5955ed3c6e8..64bf6aacf9905 100644
--- a/mlir/lib/CAPI/Dialect/Linalg.cpp
+++ b/mlir/lib/CAPI/Dialect/Linalg.cpp
@@ -147,7 +147,9 @@ mlirLinalgInferConvolutionDimensions(MlirOperation op) {
MLIR_CAPI_EXPORTED MlirLinalgConvolutionDimensions
mlirLinalgInferConvolutionDimensionsFromMaps(const MlirAffineMap *indexingMaps,
size_t numMaps) {
- if (!indexingMaps || numMaps == 0)
+ // inferConvolutionDims requires exactly 3 maps (input, filter, output);
+ // keep this check in sync with its contract
+ if (!indexingMaps || numMaps != 3)
return {};
SmallVector<AffineMap, 3> maps;
>From e877e6b991f208bab6106f76e4d6246b1690e703 Mon Sep 17 00:00:00 2001
From: Pooja Hemashekar <hemashekar at roofline.ai>
Date: Fri, 12 Jun 2026 11:11:33 +0000
Subject: [PATCH 5/5] address comments
---
mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp | 2 +-
mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp | 2 +-
2 files changed, 2 insertions(+), 2 deletions(-)
diff --git a/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp b/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp
index f3984177dba2f..e3f4988b3ab76 100644
--- a/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp
+++ b/mlir/lib/Dialect/Linalg/IR/LinalgInterfaces.cpp
@@ -748,7 +748,7 @@ getConstantsFromExprList(const SmallVector<AffineExpr, 2> &exprs) {
return vals;
}
-/// Classifies dimensions in the `linalgOp` used by a convolution
+/// Classifies dimensions in the `indexingMaps` used by a convolution
/// subcomputation, as captured by `inputExprWalker`. If
/// `allowEmptyConvolvedDims` is not set this will fail if there is not
/// at least one convolved dimension pair (output image + filter loop).
diff --git a/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp b/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp
index b5a3747660a3f..987345d7fb584 100644
--- a/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp
+++ b/mlir/unittests/Dialect/Linalg/InferConvolutionDimsTest.cpp
@@ -294,7 +294,7 @@ TEST_F(InferConvolutionDimsTest, MapsOverloadRejectsInvalidInput) {
EXPECT_TRUE(failed(inferConvolutionDims(SmallVector<AffineMap>{m, m})));
}
- // Output map is not a projected permutation: iterator types are unknowable.
+ // Output map is not a projected permutation: iterator types are unidentifiable.
{
SmallVector<AffineMap> maps = {AffineMap::get(4, 0, {d0 + d2, d1 + d3}, c),
AffineMap::get(4, 0, {d2, d3}, c),
More information about the Mlir-commits
mailing list