[RFC] Partial folding for multi-result ops

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::fold breaks 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.if and its ReplaceIfYieldWithConditionOrValue canonicalization pattern, which forwards a value that both branches of scf.if yield.

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 result i with it.
  • A Value other than the op’s own result i: the driver replaces result i with it.
  • Null, or the op’s own result i: the driver keeps result i.

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):

  1. [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.
  2. Land the core patch first: the new type, the hooks, the Operation::fold overloads, the traits and the adapters. Draft implementation.
  3. I add a dialect bit for the migration, useOpFoldResults, like the old useFoldAPI. ODS then declares the new form for the ops of that dialect. The test dialect opts in first, with test ops for every case.
  4. I teach the drivers to apply partial folds, and I update the docs.
  5. I migrate the in-tree dialects, including cir and Flang. Most bodies require no change. The same step adopts partial folds in relevant ops like the ones mentioned above.
  6. I migrate the callers of the old Operation::fold overloads, and I mark the old APIs deprecated.
  7. I add a deprecation warning for dialects that leave the bit at 0, and I post a deprecation notice.
  8. I flip the default 4 weeks after the deprecation notice.
  9. 2 weeks after the flip, I remove the bit, and the legacy code supporting the old fold signature.

Alternatives

  • OpFoldResults for 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

Thanks for looking into this.

You’re claiming this is a problem, but you’re just describing a fact that there are two APIs without explaining why it is an actual problem?

And further:

Your proposal does not address the “problem 1”, so at this point I’m confused: you’re starting your whole RFC with a problem that you don’t address? (or maybe you do, the “problem” wasn’t really well defined so I don’t know…)

This isn’t specific to graph region: you can have this situation in SSA region in unreachable code, we’re just trying for the greedy driver to avoid processing operations in unreachable code, but this is independent of how we define the properties of the folder / the fold API.

This isn’t a “work-around”: this is a blatant API violation, this shouldn’t exist. The “Canonicalization patterns do work that a fold could do” is the best symptom of the actual problem to solve:

As I understand it all (I had to read it 2/3 times to really get there), this is the actual problem to solve, everything else is noise and the proposal could be resumed to something like:

To enable multi-result ops to fold partially, we need to change the fold API to be able to express the various combination of the possible results for the fold API: the folder can optionally change the op in place or not, and also optionally fold each result. The current API cannot express it:
Here are the possible options to fix the API to be able to express this semantics:

Right now I’m missing a bit on what informations can’t be conveyed with the current API:

LogicalResult MyOp::fold(FoldAdaptor adaptor, SmallVectorImpl<OpFoldResult> &results);

Seems to me that here the LogicalResult returning “failed” can be used to express “didn’t fold” and the vector covers the individual results. Maybe we can’t clearly express “in place + some results”? The LogicalResult could express “in place” vs “not in-place” and the vector being empty or not expresses the result folding.

Yeah, you’re right. I don’t think the API asymmetry is a problem in itself. It’s true though that this is a good opportunity to bring both APIs closer (although we don’t fully cover the gap, as you call out). We can drop problem 1 from the proposal and focus on the real problem: lack of partial folding support. I’ll reframe the RFC accordingly.

TIL this is possible. Just for my own education: how does reachability affect this? Is it just that we won’t be checking it or is it actually “allowed” to break this property on unreachable code?

Yep, and that’s the first instance of this problem I ran into. AFAIK this doesn’t lead to miscompilations at the moment, as the fold itself checks whether the results have users, but it leads to difficult to predict behaviour.

Yeah, we would have 4 states that way encoded in the vector length + the LogicalResult. However, this would introduce a silent breaking API change by encoding the “in place”-ness in the LogicalResult, which may affect downstream users. I think changing the API completely to a struct with the explicit encoding is clearer (no need to “decode” from the function result + the output operand) and will force users to update their code, it won’t lead to silent bugs.


I updated the RFC dropping some details on needed changes on drivers and focusing on the actual problem (partial fold result).

This is explicitly allowed in unreachable code. It’s the case in LLVM IR and other SSA IR as well. This keeps the transformations valid locally.

We should cover this under an MLIR_EXPENSIVE_CHECK maybe? (@matthias-springer for vis)

Seems like a good fit. We already report in-place modification and op erasure that bypasses the rewriter with this mechanism. It could be as simple as maintaining a set of all op/block pointers and compare with the set after each pattern / folder application (vibe-coded prototype). I don’t have time to implement this nicely at the moment, but if somebody wants to give it try, I can review.

We can make this part of the migration if we think it’s worth it to enforce the contract. We would need to migrate the operations currently breaking it to an alternative solution (partial folding support or move work to the canonicalizer if we decide against this proposal).