Hi all,
I want to propose changing how ops fold when they do not have exactly one result.
At the moment, an op with exactly one result (variadic ops or ops with an optional result don’t apply here) use:
OpFoldResult MyOp::fold(FoldAdaptor adaptor);
Every other op uses:
LogicalResult MyOp::fold(FoldAdaptor adaptor, SmallVectorImpl<OpFoldResult> &results);
We focus on this second API in this RFC.
At the moment, on success(), the results vector must be empty (the op changed in place) or hold exactly one entry per result. By using this design, MLIR introduces a limitation: “partial folding is not supported”.
Partial folding is useful. In fact, I have found some operations that try to “work around” this limitation by breaking the API contract or moving the work to their canonicalizer:
memref.extract_strided_metadata::foldbreaks the API contract by creating constants and signaling an “in-place” change. This is “useful” as it allows for partial folding, but it is dangerous, as no driver expects this kind of behaviour.- Canonicalization patterns do work that a fold could do. This is the case of several operations like
scf.ifand itsReplaceIfYieldWithConditionOrValuecanonicalization pattern, which forwards a value that both branches ofscf.ifyield.
Proposal
Ops with exactly one result do not change. They keep OpFoldResult fold(FoldAdaptor).
Every other op returns a new type:
OpFoldResults MyOp::fold(FoldAdaptor adaptor);
OpFoldResults holds one slot per result, plus a bit that says “the fold changed the op in place”:
class [[nodiscard]] OpFoldResults {
public:
// failure
OpFoldResults() = default;
// failure
OpFoldResults(std::nullptr_t);
// success(): in place; failure(): failure
OpFoldResults(LogicalResult status);
// One slot (the op has one result at run time).
OpFoldResults(OpFoldResult replacement);
OpFoldResults(Value replacement);
OpFoldResults(Attribute replacement);
// One slot per element.
OpFoldResults(std::initializer_list<OpFoldResult> replacements);
template <typename RangeT> OpFoldResults(RangeT &&replacements);
// Intended use: build `OpFoldResults` with current operation and call `replace`.
explicit OpFoldResults(Operation *op);
void replace(Value result, OpFoldResult replacement);
void replace(unsigned resultIndex, OpFoldResult replacement);
void markModifiedInPlace(bool modified = true);
// Queries for drivers: modifiedInPlace(), replacesAny(), replacesAll(), operator[], ...
};
Semantics
Each slot i holds one of three things:
- An
Attribute: the driver materializes a constant and replaces resultiwith it. - A
Valueother than the op’s own resulti: the driver replaces resultiwith it. - Null, or the op’s own result
i: the driver keeps resulti.
Additionally:
- If the fold replaces nothing and does not mark an in-place change, it must signal failure, meaning the IR is unchanged.
- The old rules stay: a fold creates no ops, changes no IR outside the op, and returns only values that already exist.
Migration
I want to follow the path of the FoldAdaptor migration ([mlir] Add a new fold API using Generic Adaptors · llvm/llvm-project@bbfa7ef · GitHub, [mlir][tblgen] Emit deprecation warning if `kEmitRawAttributes` is used · llvm/llvm-project@d7daa63 · GitHub, [mlir] Switch default Fold API to using FoldAdaptors · llvm/llvm-project@ad48a0e · GitHub, [mlir] Complety remove old `fold` API · llvm/llvm-project@475bbea · GitHub):
- [mlir] Keep existing values when fold materialization fails by victor-eds · Pull Request #227325 · llvm/llvm-project · GitHub fixes a greedy-driver bug: after a failed constant materialization, its cleanup erases the defining ops of existing replacement values.
- Land the core patch first: the new type, the hooks, the
Operation::foldoverloads, the traits and the adapters. Draft implementation. - I add a dialect bit for the migration,
useOpFoldResults, like the olduseFoldAPI. ODS then declares the new form for the ops of that dialect. The test dialect opts in first, with test ops for every case. - I teach the drivers to apply partial folds, and I update the docs.
- I migrate the in-tree dialects, including
cirand Flang. Most bodies require no change. The same step adopts partial folds in relevant ops like the ones mentioned above. - I migrate the callers of the old
Operation::foldoverloads, and I mark the old APIs deprecated. - I add a deprecation warning for dialects that leave the bit at 0, and I post a deprecation notice.
- I flip the default 4 weeks after the deprecation notice.
- 2 weeks after the flip, I remove the bit, and the legacy code supporting the old
foldsignature.
Alternatives
OpFoldResultsfor single-result ops too. Huge splash radius for no gain at all.- Not changing the API, encode “inplace”-ness in
LogicalResult: This would silently introduce a breaking API change. Downstream projects may break. Changing the API enables a gradual migration including a stage in which the old API stays in a deprecated state, giving users time to migrate. - One hook per result (
OpFoldResult foldSum(FoldAdaptor)). It looks the most like the single-result form, but we may need to repeat some work across different results. FailureOr<SmallVector<OpFoldResult>>. It cannot express “partial and in place”.
Feedback is welcome. I will push the PR introducing the new signature today so we can talk over code. (Done: [mlir] Add OpFoldResults for partial folds of multi-result ops by victor-eds · Pull Request #227430 · llvm/llvm-project · GitHub)
Thanks,
Victor