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> ®ions) {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