20static AffineExpr getTripCountExpr(OpFoldResult lb, OpFoldResult ub,
22 ValueBoundsConstraintSet &cstr) {
23 AffineExpr lbExpr = cstr.
getExpr(lb);
24 AffineExpr ubExpr = cstr.
getExpr(ub);
25 AffineExpr stepExpr = cstr.
getExpr(step);
26 AffineExpr tripCountExpr =
27 AffineExpr(ubExpr - lbExpr).
ceilDiv(stepExpr);
31static void populateIVBounds(OpFoldResult lb, OpFoldResult ub,
32 OpFoldResult step, Value iv,
33 ValueBoundsConstraintSet &cstr) {
40 AffineExpr tripCountMinusOne =
41 getTripCountExpr(lb, ub, step, cstr) - cstr.
getExpr(1);
42 AffineExpr computedUpperBound =
43 cstr.
getExpr(lb) + AffineExpr(tripCountMinusOne * cstr.
getExpr(step));
44 cstr.
bound(iv) <= computedUpperBound;
48 :
public ValueBoundsOpInterface::ExternalModel<ForOpInterface, ForOp> {
72 static void populateIterArgBounds(scf::ForOp forOp, Value value,
73 std::optional<int64_t> dim,
74 ValueBoundsConstraintSet &cstr) {
77 if (
auto iterArg = llvm::dyn_cast<BlockArgument>(value)) {
78 iterArgIdx = iterArg.getArgNumber() - forOp.getNumInductionVars();
80 iterArgIdx = llvm::cast<OpResult>(value).getResultNumber();
83 Value yieldedValue = cast<scf::YieldOp>(forOp.getBody()->getTerminator())
84 .getOperand(iterArgIdx);
85 Value iterArg = forOp.getRegionIterArg(iterArgIdx);
86 Value initArg = forOp.getInitArgs()[iterArgIdx];
94 if (dim.has_value()) {
101 if (dim.has_value() || isa<BlockArgument>(value))
107 if (forOp.getUnsignedCmp() ||
109 forOp.getUpperBound(),
111 forOp.getLowerBound()))
117 AffineExpr tripCountExpr = getTripCountExpr(
118 forOp.getLowerBound(), forOp.getUpperBound(), forOp.getStep(), cstr);
119 AffineExpr oneIterAdvanceExpr =
122 cstr.
getExpr(initArg) + AffineExpr(tripCountExpr * oneIterAdvanceExpr);
125 void populateBoundsForIndexValue(Operation *op, Value value,
126 ValueBoundsConstraintSet &cstr)
const {
127 auto forOp = cast<ForOp>(op);
129 if (value == forOp.getInductionVar()) {
131 if (forOp.getUnsignedCmp())
133 return populateIVBounds(forOp.getLowerBound(), forOp.getUpperBound(),
134 forOp.getStep(), value, cstr);
138 populateIterArgBounds(forOp, value, std::nullopt, cstr);
141 void populateBoundsForShapedValueDim(Operation *op, Value value, int64_t dim,
142 ValueBoundsConstraintSet &cstr)
const {
143 auto forOp = cast<ForOp>(op);
145 populateIterArgBounds(forOp, value, dim, cstr);
149struct ForallOpInterface
150 :
public ValueBoundsOpInterface::ExternalModel<ForallOpInterface,
153 void populateBoundsForIndexValue(Operation *op, Value value,
154 ValueBoundsConstraintSet &cstr)
const {
155 auto forallOp = cast<ForallOp>(op);
160 auto blockArg = cast<BlockArgument>(value);
161 assert(blockArg.getArgNumber() < forallOp.getInductionVars().size() &&
162 "expected index value to be an induction var");
163 int64_t idx = blockArg.getArgNumber();
164 return populateIVBounds(forallOp.getMixedLowerBound()[idx],
165 forallOp.getMixedUpperBound()[idx],
166 forallOp.getMixedStep()[idx], value, cstr);
169 void populateBoundsForShapedValueDim(Operation *op, Value value, int64_t dim,
170 ValueBoundsConstraintSet &cstr)
const {
171 auto forallOp = cast<ForallOp>(op);
175 if (
auto iterArg = llvm::dyn_cast<BlockArgument>(value)) {
176 iterArgIdx = iterArg.getArgNumber() - forallOp.getInductionVars().size();
178 iterArgIdx = llvm::cast<OpResult>(value).getResultNumber();
183 Value outputOperand = forallOp.getOutputs()[iterArgIdx];
184 cstr.
bound(value)[dim] == cstr.
getExpr(outputOperand, dim);
189 :
public ValueBoundsOpInterface::ExternalModel<IfOpInterface, IfOp> {
191 static void populateBounds(scf::IfOp ifOp, Value value,
192 std::optional<int64_t> dim,
193 ValueBoundsConstraintSet &cstr) {
194 unsigned int resultNum = cast<OpResult>(value).getResultNumber();
195 Value thenValue = ifOp.thenYield().getResults()[resultNum];
196 Value elseValue = ifOp.elseYield().getResults()[resultNum];
198 auto boundsBuilder = cstr.
bound(value);
211 cstr.
bound(value)[*dim] >= cstr.
getExpr(thenValue, dim);
212 cstr.
bound(value)[*dim] <= cstr.
getExpr(elseValue, dim);
214 cstr.
bound(value) >= thenValue;
215 cstr.
bound(value) <= elseValue;
226 cstr.
bound(value)[*dim] >= cstr.
getExpr(elseValue, dim);
227 cstr.
bound(value)[*dim] <= cstr.
getExpr(thenValue, dim);
229 cstr.
bound(value) >= elseValue;
230 cstr.
bound(value) <= thenValue;
235 void populateBoundsForIndexValue(Operation *op, Value value,
236 ValueBoundsConstraintSet &cstr)
const {
237 populateBounds(cast<IfOp>(op), value, std::nullopt, cstr);
240 void populateBoundsForShapedValueDim(Operation *op, Value value, int64_t dim,
241 ValueBoundsConstraintSet &cstr)
const {
242 populateBounds(cast<IfOp>(op), value, dim, cstr);
253 scf::ForOp::attachInterface<scf::ForOpInterface>(*ctx);
254 scf::ForallOp::attachInterface<scf::ForallOpInterface>(*ctx);
255 scf::IfOp::attachInterface<scf::IfOpInterface>(*ctx);
AffineExpr ceilDiv(uint64_t v) const
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
MLIRContext is the top-level object for a collection of MLIR operations.
AffineExpr getExpr(Value value, std::optional< int64_t > dim=std::nullopt)
Return an expression that represents the given index-typed value or shaped value dimension.
BoundBuilder bound(Value value)
Add a bound for the given index-typed value or shaped value.
bool populateAndCompare(const Variable &lhs, ComparisonOperator cmp, const Variable &rhs)
Populate constraints for lhs/rhs (until the stop condition is met).
void registerValueBoundsOpInterfaceExternalModels(DialectRegistry ®istry)
Include the generated interface declarations.