18#define GEN_PASS_DEF_ELIMINATEVECTORMASKS
19#include "mlir/Dialect/Vector/Transforms/Passes.h.inc"
33resolveAllTrueCreateMaskOp(
IRRewriter &rewriter,
34 vector::CreateMaskOp createMaskOp,
35 std::optional<VscaleRange> vscaleRange) {
36 auto maskType = createMaskOp.getVectorType();
37 auto maskTypeDimScalableFlags = maskType.getScalableDims();
38 auto maskTypeDimSizes = maskType.getShape();
40 struct UnknownMaskDim {
52 for (
auto [i, dimSize] : llvm::enumerate(createMaskOp.getOperands())) {
55 if (maskTypeDimScalableFlags[i] || intSize < maskTypeDimSizes[i])
59 if (vscaleMultiplier < maskTypeDimSizes[i])
63 unknownDims.push_back(UnknownMaskDim{i, dimSize});
67 for (
auto [i, dimSize] : unknownDims) {
76 if (maskTypeDimScalableFlags[i])
78 FailureOr<int64_t> constantLowerBound =
81 if (
failed(constantLowerBound))
84 if (*constantLowerBound < maskTypeDimSizes[i])
91 FailureOr<ConstantOrScalableBound> dimLowerBound =
93 dimSize, {}, vscaleRange->vscaleMin, vscaleRange->vscaleMax,
97 auto dimLowerBoundSize = dimLowerBound->getSize();
98 if (
failed(dimLowerBoundSize))
100 if (dimLowerBoundSize->scalable) {
103 if (dimLowerBoundSize->baseSize < maskTypeDimSizes[i])
108 if (maskTypeDimScalableFlags[i])
111 if (dimLowerBoundSize->baseSize < maskTypeDimSizes[i])
119 auto allTrue = vector::ConstantMaskOp::create(
130 std::optional<VscaleRange> vscaleRange) {
132 if (function.isExternal())
140 function.walk([&](vector::CreateMaskOp createMaskOp) {
141 worklist.push_back(createMaskOp);
145 for (
auto mask : worklist)
146 (
void)resolveAllTrueCreateMaskOp(rewriter, mask, vscaleRange);
150struct EliminateVectorMasksPass
151 :
public impl::EliminateVectorMasksBase<EliminateVectorMasksPass> {
157 bool unset = !vscaleMin && !vscaleMax;
158 bool valid = vscaleMin && vscaleMax && vscaleMin <= vscaleMax;
161 return emitError(UnknownLoc::get(context))
162 <<
"invalid vscale range 'vscale-min="
163 <<
static_cast<unsigned>(vscaleMin)
164 <<
" vscale-max=" <<
static_cast<unsigned>(vscaleMax)
165 <<
"': expected both to be 0 (unknown), or both non-zero with "
166 "'vscale-min' <= 'vscale-max'";
169 void runOnOperation()
override {
170 std::optional<VscaleRange> vscaleRange;
171 if (vscaleMin && vscaleMax)
172 vscaleRange = VscaleRange{vscaleMin, vscaleMax};
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
MLIRContext is the top-level object for a collection of MLIR operations.
RAII guard to reset the insertion point of the builder when destroyed.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
static FailureOr< int64_t > computeConstantBound(presburger::BoundType type, const Variable &var, const StopConditionFn &stopCondition=nullptr, ValueBoundsOptions options={})
Compute a constant bound for the given variable.
std::optional< int64_t > getConstantVscaleMultiplier(Value value)
If value is a constant multiple of vector.vscale (e.g.
void eliminateVectorMasks(IRRewriter &rewriter, FunctionOpInterface function, std::optional< VscaleRange > vscaleRange={})
Split a vector.transfer operation into an in-bounds (i.e., no out-of-bounds masking) fastpath and a s...
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
static FailureOr< ConstantOrScalableBound > computeScalableBound(Value value, std::optional< int64_t > dim, unsigned vscaleMin, unsigned vscaleMax, presburger::BoundType boundType, ValueBoundsOptions options={true}, const StopConditionFn &stopCondition=nullptr)
Computes a (possibly) scalable bound for a given value.