brintos

brintos / llvm-project-archived public Read only

0
0
Text · 40.2 KiB · 7147f2c Raw
998 lines · cpp
1//===- IRAffine.cpp - Exports 'ir' module affine related bindings ---------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8 9#include <cstddef>10#include <cstdint>11#include <stdexcept>12#include <string>13#include <utility>14#include <vector>15 16#include "IRModule.h"17#include "NanobindUtils.h"18#include "mlir-c/AffineExpr.h"19#include "mlir-c/AffineMap.h"20#include "mlir-c/Bindings/Python/Interop.h" // This is expected after nanobind.21#include "mlir-c/IntegerSet.h"22#include "mlir/Bindings/Python/Nanobind.h"23#include "mlir/Support/LLVM.h"24#include "llvm/ADT/Hashing.h"25#include "llvm/ADT/SmallVector.h"26#include "llvm/ADT/StringRef.h"27#include "llvm/ADT/Twine.h"28 29namespace nb = nanobind;30using namespace mlir;31using namespace mlir::python;32 33using llvm::SmallVector;34using llvm::StringRef;35using llvm::Twine;36 37static const char kDumpDocstring[] =38    R"(Dumps a debug representation of the object to stderr.)";39 40/// Attempts to populate `result` with the content of `list` casted to the41/// appropriate type (Python and C types are provided as template arguments).42/// Throws errors in case of failure, using "action" to describe what the caller43/// was attempting to do.44template <typename PyType, typename CType>45static void pyListToVector(const nb::list &list,46                           llvm::SmallVectorImpl<CType> &result,47                           StringRef action) {48  result.reserve(nb::len(list));49  for (nb::handle item : list) {50    try {51      result.push_back(nb::cast<PyType>(item));52    } catch (nb::cast_error &err) {53      std::string msg = (llvm::Twine("Invalid expression when ") + action +54                         " (" + err.what() + ")")55                            .str();56      throw std::runtime_error(msg.c_str());57    } catch (std::runtime_error &err) {58      std::string msg = (llvm::Twine("Invalid expression (None?) when ") +59                         action + " (" + err.what() + ")")60                            .str();61      throw std::runtime_error(msg.c_str());62    }63  }64}65 66template <typename PermutationTy>67static bool isPermutation(const std::vector<PermutationTy> &permutation) {68  llvm::SmallVector<bool, 8> seen(permutation.size(), false);69  for (auto val : permutation) {70    if (val < permutation.size()) {71      if (seen[val])72        return false;73      seen[val] = true;74      continue;75    }76    return false;77  }78  return true;79}80 81namespace {82 83/// CRTP base class for Python MLIR affine expressions that subclass AffineExpr84/// and should be castable from it. Intermediate hierarchy classes can be85/// modeled by specifying BaseTy.86template <typename DerivedTy, typename BaseTy = PyAffineExpr>87class PyConcreteAffineExpr : public BaseTy {88public:89  // Derived classes must define statics for:90  //   IsAFunctionTy isaFunction91  //   const char *pyClassName92  // and redefine bindDerived.93  using ClassTy = nb::class_<DerivedTy, BaseTy>;94  using IsAFunctionTy = bool (*)(MlirAffineExpr);95 96  PyConcreteAffineExpr() = default;97  PyConcreteAffineExpr(PyMlirContextRef contextRef, MlirAffineExpr affineExpr)98      : BaseTy(std::move(contextRef), affineExpr) {}99  PyConcreteAffineExpr(PyAffineExpr &orig)100      : PyConcreteAffineExpr(orig.getContext(), castFrom(orig)) {}101 102  static MlirAffineExpr castFrom(PyAffineExpr &orig) {103    if (!DerivedTy::isaFunction(orig)) {104      auto origRepr = nb::cast<std::string>(nb::repr(nb::cast(orig)));105      throw nb::value_error((Twine("Cannot cast affine expression to ") +106                             DerivedTy::pyClassName + " (from " + origRepr +107                             ")")108                                .str()109                                .c_str());110    }111    return orig;112  }113 114  static void bind(nb::module_ &m) {115    auto cls = ClassTy(m, DerivedTy::pyClassName);116    cls.def(nb::init<PyAffineExpr &>(), nb::arg("expr"));117    cls.def_static(118        "isinstance",119        [](PyAffineExpr &otherAffineExpr) -> bool {120          return DerivedTy::isaFunction(otherAffineExpr);121        },122        nb::arg("other"));123    DerivedTy::bindDerived(cls);124  }125 126  /// Implemented by derived classes to add methods to the Python subclass.127  static void bindDerived(ClassTy &m) {}128};129 130class PyAffineConstantExpr : public PyConcreteAffineExpr<PyAffineConstantExpr> {131public:132  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsAConstant;133  static constexpr const char *pyClassName = "AffineConstantExpr";134  using PyConcreteAffineExpr::PyConcreteAffineExpr;135 136  static PyAffineConstantExpr get(intptr_t value,137                                  DefaultingPyMlirContext context) {138    MlirAffineExpr affineExpr =139        mlirAffineConstantExprGet(context->get(), static_cast<int64_t>(value));140    return PyAffineConstantExpr(context->getRef(), affineExpr);141  }142 143  static void bindDerived(ClassTy &c) {144    c.def_static("get", &PyAffineConstantExpr::get, nb::arg("value"),145                 nb::arg("context") = nb::none());146    c.def_prop_ro("value", [](PyAffineConstantExpr &self) {147      return mlirAffineConstantExprGetValue(self);148    });149  }150};151 152class PyAffineDimExpr : public PyConcreteAffineExpr<PyAffineDimExpr> {153public:154  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsADim;155  static constexpr const char *pyClassName = "AffineDimExpr";156  using PyConcreteAffineExpr::PyConcreteAffineExpr;157 158  static PyAffineDimExpr get(intptr_t pos, DefaultingPyMlirContext context) {159    MlirAffineExpr affineExpr = mlirAffineDimExprGet(context->get(), pos);160    return PyAffineDimExpr(context->getRef(), affineExpr);161  }162 163  static void bindDerived(ClassTy &c) {164    c.def_static("get", &PyAffineDimExpr::get, nb::arg("position"),165                 nb::arg("context") = nb::none());166    c.def_prop_ro("position", [](PyAffineDimExpr &self) {167      return mlirAffineDimExprGetPosition(self);168    });169  }170};171 172class PyAffineSymbolExpr : public PyConcreteAffineExpr<PyAffineSymbolExpr> {173public:174  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsASymbol;175  static constexpr const char *pyClassName = "AffineSymbolExpr";176  using PyConcreteAffineExpr::PyConcreteAffineExpr;177 178  static PyAffineSymbolExpr get(intptr_t pos, DefaultingPyMlirContext context) {179    MlirAffineExpr affineExpr = mlirAffineSymbolExprGet(context->get(), pos);180    return PyAffineSymbolExpr(context->getRef(), affineExpr);181  }182 183  static void bindDerived(ClassTy &c) {184    c.def_static("get", &PyAffineSymbolExpr::get, nb::arg("position"),185                 nb::arg("context") = nb::none());186    c.def_prop_ro("position", [](PyAffineSymbolExpr &self) {187      return mlirAffineSymbolExprGetPosition(self);188    });189  }190};191 192class PyAffineBinaryExpr : public PyConcreteAffineExpr<PyAffineBinaryExpr> {193public:194  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsABinary;195  static constexpr const char *pyClassName = "AffineBinaryExpr";196  using PyConcreteAffineExpr::PyConcreteAffineExpr;197 198  PyAffineExpr lhs() {199    MlirAffineExpr lhsExpr = mlirAffineBinaryOpExprGetLHS(get());200    return PyAffineExpr(getContext(), lhsExpr);201  }202 203  PyAffineExpr rhs() {204    MlirAffineExpr rhsExpr = mlirAffineBinaryOpExprGetRHS(get());205    return PyAffineExpr(getContext(), rhsExpr);206  }207 208  static void bindDerived(ClassTy &c) {209    c.def_prop_ro("lhs", &PyAffineBinaryExpr::lhs);210    c.def_prop_ro("rhs", &PyAffineBinaryExpr::rhs);211  }212};213 214class PyAffineAddExpr215    : public PyConcreteAffineExpr<PyAffineAddExpr, PyAffineBinaryExpr> {216public:217  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsAAdd;218  static constexpr const char *pyClassName = "AffineAddExpr";219  using PyConcreteAffineExpr::PyConcreteAffineExpr;220 221  static PyAffineAddExpr get(PyAffineExpr lhs, const PyAffineExpr &rhs) {222    MlirAffineExpr expr = mlirAffineAddExprGet(lhs, rhs);223    return PyAffineAddExpr(lhs.getContext(), expr);224  }225 226  static PyAffineAddExpr getRHSConstant(PyAffineExpr lhs, intptr_t rhs) {227    MlirAffineExpr expr = mlirAffineAddExprGet(228        lhs, mlirAffineConstantExprGet(mlirAffineExprGetContext(lhs), rhs));229    return PyAffineAddExpr(lhs.getContext(), expr);230  }231 232  static PyAffineAddExpr getLHSConstant(intptr_t lhs, PyAffineExpr rhs) {233    MlirAffineExpr expr = mlirAffineAddExprGet(234        mlirAffineConstantExprGet(mlirAffineExprGetContext(rhs), lhs), rhs);235    return PyAffineAddExpr(rhs.getContext(), expr);236  }237 238  static void bindDerived(ClassTy &c) {239    c.def_static("get", &PyAffineAddExpr::get);240  }241};242 243class PyAffineMulExpr244    : public PyConcreteAffineExpr<PyAffineMulExpr, PyAffineBinaryExpr> {245public:246  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsAMul;247  static constexpr const char *pyClassName = "AffineMulExpr";248  using PyConcreteAffineExpr::PyConcreteAffineExpr;249 250  static PyAffineMulExpr get(PyAffineExpr lhs, const PyAffineExpr &rhs) {251    MlirAffineExpr expr = mlirAffineMulExprGet(lhs, rhs);252    return PyAffineMulExpr(lhs.getContext(), expr);253  }254 255  static PyAffineMulExpr getRHSConstant(PyAffineExpr lhs, intptr_t rhs) {256    MlirAffineExpr expr = mlirAffineMulExprGet(257        lhs, mlirAffineConstantExprGet(mlirAffineExprGetContext(lhs), rhs));258    return PyAffineMulExpr(lhs.getContext(), expr);259  }260 261  static PyAffineMulExpr getLHSConstant(intptr_t lhs, PyAffineExpr rhs) {262    MlirAffineExpr expr = mlirAffineMulExprGet(263        mlirAffineConstantExprGet(mlirAffineExprGetContext(rhs), lhs), rhs);264    return PyAffineMulExpr(rhs.getContext(), expr);265  }266 267  static void bindDerived(ClassTy &c) {268    c.def_static("get", &PyAffineMulExpr::get);269  }270};271 272class PyAffineModExpr273    : public PyConcreteAffineExpr<PyAffineModExpr, PyAffineBinaryExpr> {274public:275  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsAMod;276  static constexpr const char *pyClassName = "AffineModExpr";277  using PyConcreteAffineExpr::PyConcreteAffineExpr;278 279  static PyAffineModExpr get(PyAffineExpr lhs, const PyAffineExpr &rhs) {280    MlirAffineExpr expr = mlirAffineModExprGet(lhs, rhs);281    return PyAffineModExpr(lhs.getContext(), expr);282  }283 284  static PyAffineModExpr getRHSConstant(PyAffineExpr lhs, intptr_t rhs) {285    MlirAffineExpr expr = mlirAffineModExprGet(286        lhs, mlirAffineConstantExprGet(mlirAffineExprGetContext(lhs), rhs));287    return PyAffineModExpr(lhs.getContext(), expr);288  }289 290  static PyAffineModExpr getLHSConstant(intptr_t lhs, PyAffineExpr rhs) {291    MlirAffineExpr expr = mlirAffineModExprGet(292        mlirAffineConstantExprGet(mlirAffineExprGetContext(rhs), lhs), rhs);293    return PyAffineModExpr(rhs.getContext(), expr);294  }295 296  static void bindDerived(ClassTy &c) {297    c.def_static("get", &PyAffineModExpr::get);298  }299};300 301class PyAffineFloorDivExpr302    : public PyConcreteAffineExpr<PyAffineFloorDivExpr, PyAffineBinaryExpr> {303public:304  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsAFloorDiv;305  static constexpr const char *pyClassName = "AffineFloorDivExpr";306  using PyConcreteAffineExpr::PyConcreteAffineExpr;307 308  static PyAffineFloorDivExpr get(PyAffineExpr lhs, const PyAffineExpr &rhs) {309    MlirAffineExpr expr = mlirAffineFloorDivExprGet(lhs, rhs);310    return PyAffineFloorDivExpr(lhs.getContext(), expr);311  }312 313  static PyAffineFloorDivExpr getRHSConstant(PyAffineExpr lhs, intptr_t rhs) {314    MlirAffineExpr expr = mlirAffineFloorDivExprGet(315        lhs, mlirAffineConstantExprGet(mlirAffineExprGetContext(lhs), rhs));316    return PyAffineFloorDivExpr(lhs.getContext(), expr);317  }318 319  static PyAffineFloorDivExpr getLHSConstant(intptr_t lhs, PyAffineExpr rhs) {320    MlirAffineExpr expr = mlirAffineFloorDivExprGet(321        mlirAffineConstantExprGet(mlirAffineExprGetContext(rhs), lhs), rhs);322    return PyAffineFloorDivExpr(rhs.getContext(), expr);323  }324 325  static void bindDerived(ClassTy &c) {326    c.def_static("get", &PyAffineFloorDivExpr::get);327  }328};329 330class PyAffineCeilDivExpr331    : public PyConcreteAffineExpr<PyAffineCeilDivExpr, PyAffineBinaryExpr> {332public:333  static constexpr IsAFunctionTy isaFunction = mlirAffineExprIsACeilDiv;334  static constexpr const char *pyClassName = "AffineCeilDivExpr";335  using PyConcreteAffineExpr::PyConcreteAffineExpr;336 337  static PyAffineCeilDivExpr get(PyAffineExpr lhs, const PyAffineExpr &rhs) {338    MlirAffineExpr expr = mlirAffineCeilDivExprGet(lhs, rhs);339    return PyAffineCeilDivExpr(lhs.getContext(), expr);340  }341 342  static PyAffineCeilDivExpr getRHSConstant(PyAffineExpr lhs, intptr_t rhs) {343    MlirAffineExpr expr = mlirAffineCeilDivExprGet(344        lhs, mlirAffineConstantExprGet(mlirAffineExprGetContext(lhs), rhs));345    return PyAffineCeilDivExpr(lhs.getContext(), expr);346  }347 348  static PyAffineCeilDivExpr getLHSConstant(intptr_t lhs, PyAffineExpr rhs) {349    MlirAffineExpr expr = mlirAffineCeilDivExprGet(350        mlirAffineConstantExprGet(mlirAffineExprGetContext(rhs), lhs), rhs);351    return PyAffineCeilDivExpr(rhs.getContext(), expr);352  }353 354  static void bindDerived(ClassTy &c) {355    c.def_static("get", &PyAffineCeilDivExpr::get);356  }357};358 359} // namespace360 361bool PyAffineExpr::operator==(const PyAffineExpr &other) const {362  return mlirAffineExprEqual(affineExpr, other.affineExpr);363}364 365nb::object PyAffineExpr::getCapsule() {366  return nb::steal<nb::object>(mlirPythonAffineExprToCapsule(*this));367}368 369PyAffineExpr PyAffineExpr::createFromCapsule(const nb::object &capsule) {370  MlirAffineExpr rawAffineExpr = mlirPythonCapsuleToAffineExpr(capsule.ptr());371  if (mlirAffineExprIsNull(rawAffineExpr))372    throw nb::python_error();373  return PyAffineExpr(374      PyMlirContext::forContext(mlirAffineExprGetContext(rawAffineExpr)),375      rawAffineExpr);376}377 378//------------------------------------------------------------------------------379// PyAffineMap and utilities.380//------------------------------------------------------------------------------381namespace {382 383/// A list of expressions contained in an affine map. Internally these are384/// stored as a consecutive array leading to inexpensive random access. Both385/// the map and the expression are owned by the context so we need not bother386/// with lifetime extension.387class PyAffineMapExprList388    : public Sliceable<PyAffineMapExprList, PyAffineExpr> {389public:390  static constexpr const char *pyClassName = "AffineExprList";391 392  PyAffineMapExprList(const PyAffineMap &map, intptr_t startIndex = 0,393                      intptr_t length = -1, intptr_t step = 1)394      : Sliceable(startIndex,395                  length == -1 ? mlirAffineMapGetNumResults(map) : length,396                  step),397        affineMap(map) {}398 399private:400  /// Give the parent CRTP class access to hook implementations below.401  friend class Sliceable<PyAffineMapExprList, PyAffineExpr>;402 403  intptr_t getRawNumElements() { return mlirAffineMapGetNumResults(affineMap); }404 405  PyAffineExpr getRawElement(intptr_t pos) {406    return PyAffineExpr(affineMap.getContext(),407                        mlirAffineMapGetResult(affineMap, pos));408  }409 410  PyAffineMapExprList slice(intptr_t startIndex, intptr_t length,411                            intptr_t step) {412    return PyAffineMapExprList(affineMap, startIndex, length, step);413  }414 415  PyAffineMap affineMap;416};417} // namespace418 419bool PyAffineMap::operator==(const PyAffineMap &other) const {420  return mlirAffineMapEqual(affineMap, other.affineMap);421}422 423nb::object PyAffineMap::getCapsule() {424  return nb::steal<nb::object>(mlirPythonAffineMapToCapsule(*this));425}426 427PyAffineMap PyAffineMap::createFromCapsule(const nb::object &capsule) {428  MlirAffineMap rawAffineMap = mlirPythonCapsuleToAffineMap(capsule.ptr());429  if (mlirAffineMapIsNull(rawAffineMap))430    throw nb::python_error();431  return PyAffineMap(432      PyMlirContext::forContext(mlirAffineMapGetContext(rawAffineMap)),433      rawAffineMap);434}435 436//------------------------------------------------------------------------------437// PyIntegerSet and utilities.438//------------------------------------------------------------------------------439namespace {440 441class PyIntegerSetConstraint {442public:443  PyIntegerSetConstraint(PyIntegerSet set, intptr_t pos)444      : set(std::move(set)), pos(pos) {}445 446  PyAffineExpr getExpr() {447    return PyAffineExpr(set.getContext(),448                        mlirIntegerSetGetConstraint(set, pos));449  }450 451  bool isEq() { return mlirIntegerSetIsConstraintEq(set, pos); }452 453  static void bind(nb::module_ &m) {454    nb::class_<PyIntegerSetConstraint>(m, "IntegerSetConstraint")455        .def_prop_ro("expr", &PyIntegerSetConstraint::getExpr)456        .def_prop_ro("is_eq", &PyIntegerSetConstraint::isEq);457  }458 459private:460  PyIntegerSet set;461  intptr_t pos;462};463 464class PyIntegerSetConstraintList465    : public Sliceable<PyIntegerSetConstraintList, PyIntegerSetConstraint> {466public:467  static constexpr const char *pyClassName = "IntegerSetConstraintList";468 469  PyIntegerSetConstraintList(const PyIntegerSet &set, intptr_t startIndex = 0,470                             intptr_t length = -1, intptr_t step = 1)471      : Sliceable(startIndex,472                  length == -1 ? mlirIntegerSetGetNumConstraints(set) : length,473                  step),474        set(set) {}475 476private:477  /// Give the parent CRTP class access to hook implementations below.478  friend class Sliceable<PyIntegerSetConstraintList, PyIntegerSetConstraint>;479 480  intptr_t getRawNumElements() { return mlirIntegerSetGetNumConstraints(set); }481 482  PyIntegerSetConstraint getRawElement(intptr_t pos) {483    return PyIntegerSetConstraint(set, pos);484  }485 486  PyIntegerSetConstraintList slice(intptr_t startIndex, intptr_t length,487                                   intptr_t step) {488    return PyIntegerSetConstraintList(set, startIndex, length, step);489  }490 491  PyIntegerSet set;492};493} // namespace494 495bool PyIntegerSet::operator==(const PyIntegerSet &other) const {496  return mlirIntegerSetEqual(integerSet, other.integerSet);497}498 499nb::object PyIntegerSet::getCapsule() {500  return nb::steal<nb::object>(mlirPythonIntegerSetToCapsule(*this));501}502 503PyIntegerSet PyIntegerSet::createFromCapsule(const nb::object &capsule) {504  MlirIntegerSet rawIntegerSet = mlirPythonCapsuleToIntegerSet(capsule.ptr());505  if (mlirIntegerSetIsNull(rawIntegerSet))506    throw nb::python_error();507  return PyIntegerSet(508      PyMlirContext::forContext(mlirIntegerSetGetContext(rawIntegerSet)),509      rawIntegerSet);510}511 512void mlir::python::populateIRAffine(nb::module_ &m) {513  //----------------------------------------------------------------------------514  // Mapping of PyAffineExpr and derived classes.515  //----------------------------------------------------------------------------516  nb::class_<PyAffineExpr>(m, "AffineExpr")517      .def_prop_ro(MLIR_PYTHON_CAPI_PTR_ATTR, &PyAffineExpr::getCapsule)518      .def(MLIR_PYTHON_CAPI_FACTORY_ATTR, &PyAffineExpr::createFromCapsule)519      .def("__add__", &PyAffineAddExpr::get)520      .def("__add__", &PyAffineAddExpr::getRHSConstant)521      .def("__radd__", &PyAffineAddExpr::getRHSConstant)522      .def("__mul__", &PyAffineMulExpr::get)523      .def("__mul__", &PyAffineMulExpr::getRHSConstant)524      .def("__rmul__", &PyAffineMulExpr::getRHSConstant)525      .def("__mod__", &PyAffineModExpr::get)526      .def("__mod__", &PyAffineModExpr::getRHSConstant)527      .def("__rmod__",528           [](PyAffineExpr &self, intptr_t other) {529             return PyAffineModExpr::get(530                 PyAffineConstantExpr::get(other, *self.getContext().get()),531                 self);532           })533      .def("__sub__",534           [](PyAffineExpr &self, PyAffineExpr &other) {535             auto negOne =536                 PyAffineConstantExpr::get(-1, *self.getContext().get());537             return PyAffineAddExpr::get(self,538                                         PyAffineMulExpr::get(negOne, other));539           })540      .def("__sub__",541           [](PyAffineExpr &self, intptr_t other) {542             return PyAffineAddExpr::get(543                 self,544                 PyAffineConstantExpr::get(-other, *self.getContext().get()));545           })546      .def("__rsub__",547           [](PyAffineExpr &self, intptr_t other) {548             return PyAffineAddExpr::getLHSConstant(549                 other, PyAffineMulExpr::getLHSConstant(-1, self));550           })551      .def("__eq__", [](PyAffineExpr &self,552                        PyAffineExpr &other) { return self == other; })553      .def("__eq__",554           [](PyAffineExpr &self, nb::object &other) { return false; })555      .def("__str__",556           [](PyAffineExpr &self) {557             PyPrintAccumulator printAccum;558             mlirAffineExprPrint(self, printAccum.getCallback(),559                                 printAccum.getUserData());560             return printAccum.join();561           })562      .def("__repr__",563           [](PyAffineExpr &self) {564             PyPrintAccumulator printAccum;565             printAccum.parts.append("AffineExpr(");566             mlirAffineExprPrint(self, printAccum.getCallback(),567                                 printAccum.getUserData());568             printAccum.parts.append(")");569             return printAccum.join();570           })571      .def("__hash__",572           [](PyAffineExpr &self) {573             return static_cast<size_t>(llvm::hash_value(self.get().ptr));574           })575      .def_prop_ro(576          "context",577          [](PyAffineExpr &self) -> nb::typed<nb::object, PyMlirContext> {578            return self.getContext().getObject();579          })580      .def("compose",581           [](PyAffineExpr &self, PyAffineMap &other) {582             return PyAffineExpr(self.getContext(),583                                 mlirAffineExprCompose(self, other));584           })585      .def(586          "shift_dims",587          [](PyAffineExpr &self, uint32_t numDims, uint32_t shift,588             uint32_t offset) {589            return PyAffineExpr(590                self.getContext(),591                mlirAffineExprShiftDims(self, numDims, shift, offset));592          },593          nb::arg("num_dims"), nb::arg("shift"), nb::arg("offset") = 0)594      .def(595          "shift_symbols",596          [](PyAffineExpr &self, uint32_t numSymbols, uint32_t shift,597             uint32_t offset) {598            return PyAffineExpr(599                self.getContext(),600                mlirAffineExprShiftSymbols(self, numSymbols, shift, offset));601          },602          nb::arg("num_symbols"), nb::arg("shift"), nb::arg("offset") = 0)603      .def_static(604          "simplify_affine_expr",605          [](PyAffineExpr &self, uint32_t numDims, uint32_t numSymbols) {606            return PyAffineExpr(607                self.getContext(),608                mlirSimplifyAffineExpr(self, numDims, numSymbols));609          },610          nb::arg("expr"), nb::arg("num_dims"), nb::arg("num_symbols"),611          "Simplify an affine expression by flattening and some amount of "612          "simple analysis.")613      .def_static(614          "get_add", &PyAffineAddExpr::get,615          "Gets an affine expression containing a sum of two expressions.")616      .def_static("get_add", &PyAffineAddExpr::getLHSConstant,617                  "Gets an affine expression containing a sum of a constant "618                  "and another expression.")619      .def_static("get_add", &PyAffineAddExpr::getRHSConstant,620                  "Gets an affine expression containing a sum of an expression "621                  "and a constant.")622      .def_static(623          "get_mul", &PyAffineMulExpr::get,624          "Gets an affine expression containing a product of two expressions.")625      .def_static("get_mul", &PyAffineMulExpr::getLHSConstant,626                  "Gets an affine expression containing a product of a "627                  "constant and another expression.")628      .def_static("get_mul", &PyAffineMulExpr::getRHSConstant,629                  "Gets an affine expression containing a product of an "630                  "expression and a constant.")631      .def_static("get_mod", &PyAffineModExpr::get,632                  "Gets an affine expression containing the modulo of dividing "633                  "one expression by another.")634      .def_static("get_mod", &PyAffineModExpr::getLHSConstant,635                  "Gets a semi-affine expression containing the modulo of "636                  "dividing a constant by an expression.")637      .def_static("get_mod", &PyAffineModExpr::getRHSConstant,638                  "Gets an affine expression containing the module of dividing"639                  "an expression by a constant.")640      .def_static("get_floor_div", &PyAffineFloorDivExpr::get,641                  "Gets an affine expression containing the rounded-down "642                  "result of dividing one expression by another.")643      .def_static("get_floor_div", &PyAffineFloorDivExpr::getLHSConstant,644                  "Gets a semi-affine expression containing the rounded-down "645                  "result of dividing a constant by an expression.")646      .def_static("get_floor_div", &PyAffineFloorDivExpr::getRHSConstant,647                  "Gets an affine expression containing the rounded-down "648                  "result of dividing an expression by a constant.")649      .def_static("get_ceil_div", &PyAffineCeilDivExpr::get,650                  "Gets an affine expression containing the rounded-up result "651                  "of dividing one expression by another.")652      .def_static("get_ceil_div", &PyAffineCeilDivExpr::getLHSConstant,653                  "Gets a semi-affine expression containing the rounded-up "654                  "result of dividing a constant by an expression.")655      .def_static("get_ceil_div", &PyAffineCeilDivExpr::getRHSConstant,656                  "Gets an affine expression containing the rounded-up result "657                  "of dividing an expression by a constant.")658      .def_static("get_constant", &PyAffineConstantExpr::get, nb::arg("value"),659                  nb::arg("context") = nb::none(),660                  "Gets a constant affine expression with the given value.")661      .def_static(662          "get_dim", &PyAffineDimExpr::get, nb::arg("position"),663          nb::arg("context") = nb::none(),664          "Gets an affine expression of a dimension at the given position.")665      .def_static(666          "get_symbol", &PyAffineSymbolExpr::get, nb::arg("position"),667          nb::arg("context") = nb::none(),668          "Gets an affine expression of a symbol at the given position.")669      .def(670          "dump", [](PyAffineExpr &self) { mlirAffineExprDump(self); },671          kDumpDocstring);672  PyAffineConstantExpr::bind(m);673  PyAffineDimExpr::bind(m);674  PyAffineSymbolExpr::bind(m);675  PyAffineBinaryExpr::bind(m);676  PyAffineAddExpr::bind(m);677  PyAffineMulExpr::bind(m);678  PyAffineModExpr::bind(m);679  PyAffineFloorDivExpr::bind(m);680  PyAffineCeilDivExpr::bind(m);681 682  //----------------------------------------------------------------------------683  // Mapping of PyAffineMap.684  //----------------------------------------------------------------------------685  nb::class_<PyAffineMap>(m, "AffineMap")686      .def_prop_ro(MLIR_PYTHON_CAPI_PTR_ATTR, &PyAffineMap::getCapsule)687      .def(MLIR_PYTHON_CAPI_FACTORY_ATTR, &PyAffineMap::createFromCapsule)688      .def("__eq__",689           [](PyAffineMap &self, PyAffineMap &other) { return self == other; })690      .def("__eq__", [](PyAffineMap &self, nb::object &other) { return false; })691      .def("__str__",692           [](PyAffineMap &self) {693             PyPrintAccumulator printAccum;694             mlirAffineMapPrint(self, printAccum.getCallback(),695                                printAccum.getUserData());696             return printAccum.join();697           })698      .def("__repr__",699           [](PyAffineMap &self) {700             PyPrintAccumulator printAccum;701             printAccum.parts.append("AffineMap(");702             mlirAffineMapPrint(self, printAccum.getCallback(),703                                printAccum.getUserData());704             printAccum.parts.append(")");705             return printAccum.join();706           })707      .def("__hash__",708           [](PyAffineMap &self) {709             return static_cast<size_t>(llvm::hash_value(self.get().ptr));710           })711      .def_static(712          "compress_unused_symbols",713          [](const nb::list &affineMaps, DefaultingPyMlirContext context) {714            SmallVector<MlirAffineMap> maps;715            pyListToVector<PyAffineMap, MlirAffineMap>(716                affineMaps, maps, "attempting to create an AffineMap");717            std::vector<MlirAffineMap> compressed(affineMaps.size());718            auto populate = [](void *result, intptr_t idx, MlirAffineMap m) {719              static_cast<MlirAffineMap *>(result)[idx] = (m);720            };721            mlirAffineMapCompressUnusedSymbols(maps.data(), maps.size(),722                                               compressed.data(), populate);723            std::vector<PyAffineMap> res;724            res.reserve(compressed.size());725            for (auto m : compressed)726              res.emplace_back(context->getRef(), m);727            return res;728          })729      .def_prop_ro(730          "context",731          [](PyAffineMap &self) -> nb::typed<nb::object, PyMlirContext> {732            return self.getContext().getObject();733          },734          "Context that owns the Affine Map")735      .def(736          "dump", [](PyAffineMap &self) { mlirAffineMapDump(self); },737          kDumpDocstring)738      .def_static(739          "get",740          [](intptr_t dimCount, intptr_t symbolCount, const nb::list &exprs,741             DefaultingPyMlirContext context) {742            SmallVector<MlirAffineExpr> affineExprs;743            pyListToVector<PyAffineExpr, MlirAffineExpr>(744                exprs, affineExprs, "attempting to create an AffineMap");745            MlirAffineMap map =746                mlirAffineMapGet(context->get(), dimCount, symbolCount,747                                 affineExprs.size(), affineExprs.data());748            return PyAffineMap(context->getRef(), map);749          },750          nb::arg("dim_count"), nb::arg("symbol_count"), nb::arg("exprs"),751          nb::arg("context") = nb::none(),752          "Gets a map with the given expressions as results.")753      .def_static(754          "get_constant",755          [](intptr_t value, DefaultingPyMlirContext context) {756            MlirAffineMap affineMap =757                mlirAffineMapConstantGet(context->get(), value);758            return PyAffineMap(context->getRef(), affineMap);759          },760          nb::arg("value"), nb::arg("context") = nb::none(),761          "Gets an affine map with a single constant result")762      .def_static(763          "get_empty",764          [](DefaultingPyMlirContext context) {765            MlirAffineMap affineMap = mlirAffineMapEmptyGet(context->get());766            return PyAffineMap(context->getRef(), affineMap);767          },768          nb::arg("context") = nb::none(), "Gets an empty affine map.")769      .def_static(770          "get_identity",771          [](intptr_t nDims, DefaultingPyMlirContext context) {772            MlirAffineMap affineMap =773                mlirAffineMapMultiDimIdentityGet(context->get(), nDims);774            return PyAffineMap(context->getRef(), affineMap);775          },776          nb::arg("n_dims"), nb::arg("context") = nb::none(),777          "Gets an identity map with the given number of dimensions.")778      .def_static(779          "get_minor_identity",780          [](intptr_t nDims, intptr_t nResults,781             DefaultingPyMlirContext context) {782            MlirAffineMap affineMap =783                mlirAffineMapMinorIdentityGet(context->get(), nDims, nResults);784            return PyAffineMap(context->getRef(), affineMap);785          },786          nb::arg("n_dims"), nb::arg("n_results"),787          nb::arg("context") = nb::none(),788          "Gets a minor identity map with the given number of dimensions and "789          "results.")790      .def_static(791          "get_permutation",792          [](std::vector<unsigned> permutation,793             DefaultingPyMlirContext context) {794            if (!isPermutation(permutation))795              throw std::runtime_error("Invalid permutation when attempting to "796                                       "create an AffineMap");797            MlirAffineMap affineMap = mlirAffineMapPermutationGet(798                context->get(), permutation.size(), permutation.data());799            return PyAffineMap(context->getRef(), affineMap);800          },801          nb::arg("permutation"), nb::arg("context") = nb::none(),802          "Gets an affine map that permutes its inputs.")803      .def(804          "get_submap",805          [](PyAffineMap &self, std::vector<intptr_t> &resultPos) {806            intptr_t numResults = mlirAffineMapGetNumResults(self);807            for (intptr_t pos : resultPos) {808              if (pos < 0 || pos >= numResults)809                throw nb::value_error("result position out of bounds");810            }811            MlirAffineMap affineMap = mlirAffineMapGetSubMap(812                self, resultPos.size(), resultPos.data());813            return PyAffineMap(self.getContext(), affineMap);814          },815          nb::arg("result_positions"))816      .def(817          "get_major_submap",818          [](PyAffineMap &self, intptr_t nResults) {819            if (nResults >= mlirAffineMapGetNumResults(self))820              throw nb::value_error("number of results out of bounds");821            MlirAffineMap affineMap =822                mlirAffineMapGetMajorSubMap(self, nResults);823            return PyAffineMap(self.getContext(), affineMap);824          },825          nb::arg("n_results"))826      .def(827          "get_minor_submap",828          [](PyAffineMap &self, intptr_t nResults) {829            if (nResults >= mlirAffineMapGetNumResults(self))830              throw nb::value_error("number of results out of bounds");831            MlirAffineMap affineMap =832                mlirAffineMapGetMinorSubMap(self, nResults);833            return PyAffineMap(self.getContext(), affineMap);834          },835          nb::arg("n_results"))836      .def(837          "replace",838          [](PyAffineMap &self, PyAffineExpr &expression,839             PyAffineExpr &replacement, intptr_t numResultDims,840             intptr_t numResultSyms) {841            MlirAffineMap affineMap = mlirAffineMapReplace(842                self, expression, replacement, numResultDims, numResultSyms);843            return PyAffineMap(self.getContext(), affineMap);844          },845          nb::arg("expr"), nb::arg("replacement"), nb::arg("n_result_dims"),846          nb::arg("n_result_syms"))847      .def_prop_ro(848          "is_permutation",849          [](PyAffineMap &self) { return mlirAffineMapIsPermutation(self); })850      .def_prop_ro("is_projected_permutation",851                   [](PyAffineMap &self) {852                     return mlirAffineMapIsProjectedPermutation(self);853                   })854      .def_prop_ro(855          "n_dims",856          [](PyAffineMap &self) { return mlirAffineMapGetNumDims(self); })857      .def_prop_ro(858          "n_inputs",859          [](PyAffineMap &self) { return mlirAffineMapGetNumInputs(self); })860      .def_prop_ro(861          "n_symbols",862          [](PyAffineMap &self) { return mlirAffineMapGetNumSymbols(self); })863      .def_prop_ro("results",864                   [](PyAffineMap &self) { return PyAffineMapExprList(self); });865  PyAffineMapExprList::bind(m);866 867  //----------------------------------------------------------------------------868  // Mapping of PyIntegerSet.869  //----------------------------------------------------------------------------870  nb::class_<PyIntegerSet>(m, "IntegerSet")871      .def_prop_ro(MLIR_PYTHON_CAPI_PTR_ATTR, &PyIntegerSet::getCapsule)872      .def(MLIR_PYTHON_CAPI_FACTORY_ATTR, &PyIntegerSet::createFromCapsule)873      .def("__eq__", [](PyIntegerSet &self,874                        PyIntegerSet &other) { return self == other; })875      .def("__eq__",876           [](PyIntegerSet &self, const nb::object &other) { return false; })877      .def("__str__",878           [](PyIntegerSet &self) {879             PyPrintAccumulator printAccum;880             mlirIntegerSetPrint(self, printAccum.getCallback(),881                                 printAccum.getUserData());882             return printAccum.join();883           })884      .def("__repr__",885           [](PyIntegerSet &self) {886             PyPrintAccumulator printAccum;887             printAccum.parts.append("IntegerSet(");888             mlirIntegerSetPrint(self, printAccum.getCallback(),889                                 printAccum.getUserData());890             printAccum.parts.append(")");891             return printAccum.join();892           })893      .def("__hash__",894           [](PyIntegerSet &self) {895             return static_cast<size_t>(llvm::hash_value(self.get().ptr));896           })897      .def_prop_ro(898          "context",899          [](PyIntegerSet &self) -> nb::typed<nb::object, PyMlirContext> {900            return self.getContext().getObject();901          })902      .def(903          "dump", [](PyIntegerSet &self) { mlirIntegerSetDump(self); },904          kDumpDocstring)905      .def_static(906          "get",907          [](intptr_t numDims, intptr_t numSymbols, const nb::list &exprs,908             std::vector<bool> eqFlags, DefaultingPyMlirContext context) {909            if (exprs.size() != eqFlags.size())910              throw nb::value_error(911                  "Expected the number of constraints to match "912                  "that of equality flags");913            if (exprs.size() == 0)914              throw nb::value_error("Expected non-empty list of constraints");915 916            // Copy over to a SmallVector because std::vector has a917            // specialization for booleans that packs data and does not918            // expose a `bool *`.919            SmallVector<bool, 8> flags(eqFlags.begin(), eqFlags.end());920 921            SmallVector<MlirAffineExpr> affineExprs;922            pyListToVector<PyAffineExpr>(exprs, affineExprs,923                                         "attempting to create an IntegerSet");924            MlirIntegerSet set = mlirIntegerSetGet(925                context->get(), numDims, numSymbols, exprs.size(),926                affineExprs.data(), flags.data());927            return PyIntegerSet(context->getRef(), set);928          },929          nb::arg("num_dims"), nb::arg("num_symbols"), nb::arg("exprs"),930          nb::arg("eq_flags"), nb::arg("context") = nb::none())931      .def_static(932          "get_empty",933          [](intptr_t numDims, intptr_t numSymbols,934             DefaultingPyMlirContext context) {935            MlirIntegerSet set =936                mlirIntegerSetEmptyGet(context->get(), numDims, numSymbols);937            return PyIntegerSet(context->getRef(), set);938          },939          nb::arg("num_dims"), nb::arg("num_symbols"),940          nb::arg("context") = nb::none())941      .def(942          "get_replaced",943          [](PyIntegerSet &self, const nb::list &dimExprs,944             const nb::list &symbolExprs, intptr_t numResultDims,945             intptr_t numResultSymbols) {946            if (static_cast<intptr_t>(dimExprs.size()) !=947                mlirIntegerSetGetNumDims(self))948              throw nb::value_error(949                  "Expected the number of dimension replacement expressions "950                  "to match that of dimensions");951            if (static_cast<intptr_t>(symbolExprs.size()) !=952                mlirIntegerSetGetNumSymbols(self))953              throw nb::value_error(954                  "Expected the number of symbol replacement expressions "955                  "to match that of symbols");956 957            SmallVector<MlirAffineExpr> dimAffineExprs, symbolAffineExprs;958            pyListToVector<PyAffineExpr>(959                dimExprs, dimAffineExprs,960                "attempting to create an IntegerSet by replacing dimensions");961            pyListToVector<PyAffineExpr>(962                symbolExprs, symbolAffineExprs,963                "attempting to create an IntegerSet by replacing symbols");964            MlirIntegerSet set = mlirIntegerSetReplaceGet(965                self, dimAffineExprs.data(), symbolAffineExprs.data(),966                numResultDims, numResultSymbols);967            return PyIntegerSet(self.getContext(), set);968          },969          nb::arg("dim_exprs"), nb::arg("symbol_exprs"),970          nb::arg("num_result_dims"), nb::arg("num_result_symbols"))971      .def_prop_ro("is_canonical_empty",972                   [](PyIntegerSet &self) {973                     return mlirIntegerSetIsCanonicalEmpty(self);974                   })975      .def_prop_ro(976          "n_dims",977          [](PyIntegerSet &self) { return mlirIntegerSetGetNumDims(self); })978      .def_prop_ro(979          "n_symbols",980          [](PyIntegerSet &self) { return mlirIntegerSetGetNumSymbols(self); })981      .def_prop_ro(982          "n_inputs",983          [](PyIntegerSet &self) { return mlirIntegerSetGetNumInputs(self); })984      .def_prop_ro("n_equalities",985                   [](PyIntegerSet &self) {986                     return mlirIntegerSetGetNumEqualities(self);987                   })988      .def_prop_ro("n_inequalities",989                   [](PyIntegerSet &self) {990                     return mlirIntegerSetGetNumInequalities(self);991                   })992      .def_prop_ro("constraints", [](PyIntegerSet &self) {993        return PyIntegerSetConstraintList(self);994      });995  PyIntegerSetConstraint::bind(m);996  PyIntegerSetConstraintList::bind(m);997}998