MLIR 24.0.0git
VectorMaskElimination.cpp
Go to the documentation of this file.
1//===- VectorMaskElimination.cpp - Eliminate Vector Masks -----------------===//
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-exception
6//
7//===----------------------------------------------------------------------===//
8
14
15namespace mlir {
16namespace vector {
17
18#define GEN_PASS_DEF_ELIMINATEVECTORMASKS
19#include "mlir/Dialect/Vector/Transforms/Passes.h.inc"
20
21} // namespace vector
22} // namespace mlir
23
24using namespace mlir;
25using namespace mlir::vector;
26namespace {
27
28/// Attempts to resolve a CreateMaskOp to an all-true constant mask. All-true
29/// masks can then be eliminated by simple folds. `vscaleRange` is required to
30/// reason about scalable dimensions; without it only fixed-size dimensions can
31/// be proven all-true.
32LogicalResult
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();
39
40 struct UnknownMaskDim {
41 size_t position;
42 Value dimSize;
43 };
44
45 // Loop over the CreateMaskOp operands and collect unknown dims (i.e. dims
46 // that are not obviously constant). If any constant dimension is not all-true
47 // bail out early (as this transform only trying to resolve all-true masks).
48 // This avoids doing value-bounds anaylis in cases like:
49 // `%mask = vector.create_mask %dynamicValue, %c2 : vector<8x4xi1>`
50 // ...where it is known the mask is not all-true by looking at `%c2`.
52 for (auto [i, dimSize] : llvm::enumerate(createMaskOp.getOperands())) {
53 if (auto intSize = getConstantIntValue(dimSize)) {
54 // Mask not all-true for this dim.
55 if (maskTypeDimScalableFlags[i] || intSize < maskTypeDimSizes[i])
56 return failure();
57 } else if (auto vscaleMultiplier = getConstantVscaleMultiplier(dimSize)) {
58 // Mask not all-true for this dim.
59 if (vscaleMultiplier < maskTypeDimSizes[i])
60 return failure();
61 } else {
62 // Unknown (without further analysis).
63 unknownDims.push_back(UnknownMaskDim{i, dimSize});
64 }
65 }
66
67 for (auto [i, dimSize] : unknownDims) {
68 // Compute the lower bound for the unknown dimension (i.e. the smallest
69 // value it could be).
70
71 // Fixed-width case: without a `vscale` range the bound is a plain constant,
72 // which can only prove a fixed-size dimension all-true.
73 if (!vscaleRange) {
74 // A constant bound cannot prove a scalable dim, whose runtime size is
75 // `vscale` times the size in the type. Checked first to skip the query.
76 if (maskTypeDimScalableFlags[i])
77 return failure();
78 FailureOr<int64_t> constantLowerBound =
81 if (failed(constantLowerBound))
82 return failure();
83 // If LB < the mask dim size then this dim is not all-true.
84 if (*constantLowerBound < maskTypeDimSizes[i])
85 return failure();
86 continue;
87 }
88
89 // Scalable case: with a `vscale` range the bound has the form
90 // `base + n * vscale`, which can prove either kind of dimension all-true.
91 FailureOr<ConstantOrScalableBound> dimLowerBound =
93 dimSize, {}, vscaleRange->vscaleMin, vscaleRange->vscaleMax,
95 if (failed(dimLowerBound))
96 return failure();
97 auto dimLowerBoundSize = dimLowerBound->getSize();
98 if (failed(dimLowerBoundSize))
99 return failure();
100 if (dimLowerBoundSize->scalable) {
101 // 1. The lower bound, LB, is scalable. If LB is < the mask dim size then
102 // this dim is not all-true.
103 if (dimLowerBoundSize->baseSize < maskTypeDimSizes[i])
104 return failure();
105 } else {
106 // 2. The lower bound, LB, is a constant.
107 // - If the mask dim size is scalable then this dim is not all-true.
108 if (maskTypeDimScalableFlags[i])
109 return failure();
110 // - If LB < the _fixed-size_ mask dim size then this dim is not all-true.
111 if (dimLowerBoundSize->baseSize < maskTypeDimSizes[i])
112 return failure();
113 }
114 }
115
116 // Replace createMaskOp with an all-true constant. This should result in the
117 // mask being removed in most cases (as xfer ops + vector.mask have folds to
118 // remove all-true masks).
119 auto allTrue = vector::ConstantMaskOp::create(
120 rewriter, createMaskOp.getLoc(), maskType, ConstantMaskKind::AllTrue);
121 rewriter.replaceAllUsesWith(createMaskOp, allTrue);
122 return success();
123}
124
125} // namespace
126
127namespace mlir::vector {
128
129void eliminateVectorMasks(IRRewriter &rewriter, FunctionOpInterface function,
130 std::optional<VscaleRange> vscaleRange) {
131 // Early exit for functions without a body.
132 if (function.isExternal())
133 return;
134
135 OpBuilder::InsertionGuard g(rewriter);
136
137 // Build worklist so we can safely insert new ops in
138 // `resolveAllTrueCreateMaskOp()`.
140 function.walk([&](vector::CreateMaskOp createMaskOp) {
141 worklist.push_back(createMaskOp);
142 });
143
144 rewriter.setInsertionPointToStart(&function.front());
145 for (auto mask : worklist)
146 (void)resolveAllTrueCreateMaskOp(rewriter, mask, vscaleRange);
147}
148
149namespace {
150struct EliminateVectorMasksPass
151 : public impl::EliminateVectorMasksBase<EliminateVectorMasksPass> {
152 using Base::Base;
153
154 // Checked here rather than in runOnOperation so that a bad range is reported
155 // once, not once per function.
156 LogicalResult initialize(MLIRContext *context) override {
157 bool unset = !vscaleMin && !vscaleMax;
158 bool valid = vscaleMin && vscaleMax && vscaleMin <= vscaleMax;
159 if (unset || valid)
160 return success();
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'";
167 }
168
169 void runOnOperation() override {
170 std::optional<VscaleRange> vscaleRange;
171 if (vscaleMin && vscaleMax)
172 vscaleRange = VscaleRange{vscaleMin, vscaleMax};
173
174 IRRewriter rewriter(&getContext());
175 eliminateVectorMasks(rewriter, getOperation(), vscaleRange);
176 }
177};
178} // namespace
179
180} // namespace mlir::vector
return success()
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
b getContext())
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.
Definition MLIRContext.h:63
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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.