brintos

brintos / llvm-project-archived public Read only

0
0
Text · 9.4 KiB · f727118 Raw
245 lines · cpp
1//===- TuneExtensionOps.cpp - Tune extension for the Transform dialect ----===//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 "mlir/Dialect/Transform/IR/TransformOps.h"10#include "mlir/Dialect/Transform/Interfaces/TransformInterfaces.h"11#include "mlir/IR/OpImplementation.h"12#include "mlir/Interfaces/ControlFlowInterfaces.h"13#include "llvm/Support/Debug.h"14 15#include "mlir/Dialect/Transform/TuneExtension/TuneExtensionOps.h"16 17using namespace mlir;18 19static ParseResult parseAlternativesOpSelectedRegion(20    OpAsmParser &parser, IntegerAttr &selectedRegionAttr,21    std::optional<OpAsmParser::UnresolvedOperand> &selectedRegionParam);22 23static void printAlternativesOpSelectedRegion(OpAsmPrinter &printer,24                                              Operation *op,25                                              IntegerAttr selectedRegionAttr,26                                              Value selectedRegionParam);27 28#define GET_OP_CLASSES29#include "mlir/Dialect/Transform/TuneExtension/TuneExtensionOps.cpp.inc"30 31#define DEBUG_TYPE "transform-tune"32#define DBGS() (llvm::dbgs() << "[" DEBUG_TYPE "] ")33 34//===----------------------------------------------------------------------===//35// KnobOp36//===----------------------------------------------------------------------===//37 38void transform::tune::KnobOp::getEffects(39    SmallVectorImpl<MemoryEffects::EffectInstance> &effects) {40  producesHandle(getOperation()->getOpResults(), effects);41  onlyReadsPayload(effects);42}43 44DiagnosedSilenceableFailure45transform::tune::KnobOp::apply(transform::TransformRewriter &rewriter,46                               transform::TransformResults &results,47                               transform::TransformState &state) {48  if (getSelected()) {49    results.setParams(llvm::cast<OpResult>(getResult()), *getSelected());50    return DiagnosedSilenceableFailure::success();51  }52 53  return emitDefiniteFailure()54         << "non-deterministic choice " << getName()55         << " is only resolved through providing a `selected` attr";56}57 58LogicalResult transform::tune::KnobOp::verify() {59  if (auto selected = getSelected()) {60    if (auto optionsArray = dyn_cast<ArrayAttr>(getOptions())) {61      if (!llvm::is_contained(optionsArray, selected))62        return emitOpError("provided `selected` attribute is not an element of "63                           "`options` array of attributes");64    } else65      LLVM_DEBUG(DBGS() << "cannot verify `selected` attribute " << selected66                        << " is an element of `options` attribute "67                        << getOptions());68  }69 70  return success();71}72 73//===----------------------------------------------------------------------===//74// AlternativesOp75//===----------------------------------------------------------------------===//76 77static ParseResult parseAlternativesOpSelectedRegion(78    OpAsmParser &parser, IntegerAttr &selectedRegionAttr,79    std::optional<OpAsmParser::UnresolvedOperand> &selectedRegionParam) {80  size_t selectedRegionIdx;81  OptionalParseResult attrParseRes =82      parser.parseOptionalInteger(selectedRegionIdx);83  if (attrParseRes.has_value()) {84    if (failed(*attrParseRes))85      return failure();86 87    selectedRegionAttr = parser.getBuilder().getIndexAttr(selectedRegionIdx);88    return success();89  }90 91  OpAsmParser::UnresolvedOperand param;92  auto paramParseRes = parser.parseOptionalOperand(param);93  if (paramParseRes.has_value()) {94    if (failed(*paramParseRes))95      return failure();96 97    selectedRegionParam = param;98    return success();99  }100 101  return parser.emitError(parser.getCurrentLocation())102         << "expected either an integer attribute or a transform.param operand";103}104 105static void printAlternativesOpSelectedRegion(OpAsmPrinter &printer,106                                              Operation *op,107                                              IntegerAttr selectedRegionAttr,108                                              Value selectedRegionParam) {109  if (selectedRegionAttr)110    printer << selectedRegionAttr.getValue();111  if (selectedRegionParam)112    printer << selectedRegionParam;113}114 115OperandRange transform::tune::AlternativesOp::getEntrySuccessorOperands(116    RegionSuccessor successor) {117  // No operands will be forwarded to the region(s).118  return getOperands().slice(0, 0);119}120 121void transform::tune::AlternativesOp::getSuccessorRegions(122    RegionBranchPoint point, SmallVectorImpl<RegionSuccessor> &regions) {123  if (point.isParent())124    if (auto selectedRegionIdx = getSelectedRegionAttr())125      regions.emplace_back(126          &getAlternatives()[selectedRegionIdx->getSExtValue()],127          Block::BlockArgListType());128    else129      for (Region &alternative : getAlternatives())130        regions.emplace_back(&alternative, Block::BlockArgListType());131  else132    regions.emplace_back(getOperation(), getOperation()->getResults());133}134 135void transform::tune::AlternativesOp::getRegionInvocationBounds(136    ArrayRef<Attribute> operands, SmallVectorImpl<InvocationBounds> &bounds) {137  (void)operands;138  bounds.reserve(getNumRegions());139 140  if (auto selectedRegionIdx = getSelectedRegionAttr()) {141    bounds.resize(getNumRegions(), InvocationBounds(0, 0));142    bounds[selectedRegionIdx->getSExtValue()] = InvocationBounds(1, 1);143  } else {144    bounds.resize(getNumRegions(), InvocationBounds(0, 1));145  }146}147 148void transform::tune::AlternativesOp::getEffects(149    SmallVectorImpl<MemoryEffects::EffectInstance> &effects) {150  onlyReadsHandle(getSelectedRegionParamMutable(), effects);151  producesHandle(getOperation()->getOpResults(), effects);152  // TODO: should effects from regions be forwarded?153}154 155DiagnosedSilenceableFailure156transform::tune::AlternativesOp::apply(transform::TransformRewriter &rewriter,157                                       transform::TransformResults &results,158                                       transform::TransformState &state) {159  std::optional<size_t> selectedRegionIdx;160 161  if (auto selectedRegionAttr = getSelectedRegionAttr())162    selectedRegionIdx = selectedRegionAttr->getSExtValue();163 164  if (Value selectedRegionParam = getSelectedRegionParam()) {165    ArrayRef<Attribute> associatedAttrs = state.getParams(selectedRegionParam);166    IntegerAttr selectedRegionAttr;167    if (associatedAttrs.size() != 1 ||168        !(selectedRegionAttr = dyn_cast<IntegerAttr>(associatedAttrs[0])))169      return emitDefiniteFailure()170             << "param should hold exactly one integer attribute, got: "171             << associatedAttrs[0];172    selectedRegionIdx = selectedRegionAttr.getValue().getSExtValue();173  }174 175  if (!selectedRegionIdx)176    return emitDefiniteFailure() << "non-deterministic choice " << getName()177                                 << " is only resolved through providing a "178                                    "`selected_region` attr/param";179 180  if (*selectedRegionIdx < 0 || *selectedRegionIdx >= getNumRegions())181    return emitDefiniteFailure()182           << "'selected_region' attribute/param specifies region at index "183           << *selectedRegionIdx << " while op has only " << getNumRegions()184           << " regions";185 186  Region &selectedRegion = getRegion(*selectedRegionIdx);187  auto scope = state.make_region_scope(selectedRegion);188  Block &block = selectedRegion.front();189  // Apply the region's ops one by one.190  for (Operation &transform : block.without_terminator()) {191    DiagnosedSilenceableFailure result =192        state.applyTransform(cast<transform::TransformOpInterface>(transform));193    if (result.isDefiniteFailure())194      return result;195 196    if (result.isSilenceableFailure()) {197      for (const auto &res : getResults())198        results.set(res, {});199      return result;200    }201  }202  // Forward the operation mapping for values yielded from the region to the203  // values produced by the alternatives op.204  transform::detail::forwardTerminatorOperands(&block, state, results);205  return DiagnosedSilenceableFailure::success();206}207 208LogicalResult transform::tune::AlternativesOp::verify() {209  for (auto *region : getRegions()) {210    auto yieldTerminator =211        llvm::dyn_cast_if_present<transform::YieldOp>(region->front().back());212    if (!yieldTerminator)213      return emitOpError() << "expected '"214                           << transform::YieldOp::getOperationName()215                           << "' as terminator";216 217    if (yieldTerminator->getNumOperands() != getNumResults())218      return yieldTerminator.emitOpError()219             << "expected terminator to have as many operands as the parent op "220                "has results";221 222    for (auto [i, operandType, resultType] : llvm::zip_equal(223             llvm::seq<unsigned>(0, yieldTerminator->getNumOperands()),224             yieldTerminator->getOperands().getType(), getResultTypes())) {225      if (operandType == resultType)226        continue;227      return yieldTerminator.emitOpError()228             << "the type of the terminator operand #" << i229             << " must match the type of the corresponding parent op result ("230             << operandType << " vs " << resultType << ")";231    }232  }233 234  if (auto selectedRegionAttr = getSelectedRegionAttr()) {235    size_t regionIdx = selectedRegionAttr->getSExtValue();236    if (regionIdx < 0 || regionIdx >= getNumRegions())237      return emitOpError()238             << "'selected_region' attribute specifies region at index "239             << regionIdx << " while op has only " << getNumRegions()240             << " regions";241  }242 243  return success();244}245