diff --git a/src/passes/TupleOptimization.cpp b/src/passes/TupleOptimization.cpp index 0a9584e09fd..4aac390e42f 100644 --- a/src/passes/TupleOptimization.cpp +++ b/src/passes/TupleOptimization.cpp @@ -44,6 +44,7 @@ // definitely worth lowering. // +#include #include #include #include @@ -51,6 +52,74 @@ namespace wasm { +namespace { + +// Helper class to analyze interferences between tuple.make operands and their +// target locals. When lowering a tuple local.set to individual local.sets: +// +// (local.set $t (tuple.make op0 op1 ... opN-1)) +// +// into: +// +// (local.set $t0 op0) +// (local.set $t1 op1) +// ... +// +// an operand op_i cannot be written directly to target local $ti if doing so +// interferes with the evaluation of any subsequent operand op_j (j > i). An +// interference occurs if: +// 1. op_j reads $ti (op_j would see the new value of $ti instead of the old +// value). +// 2. op_j writes $ti (op_j's write would be overwritten by op_i's value later +// in the original code, but would overwrite op_i's value here). +// 3. op_j transfers control flow (e.g. branches to an enclosing block or +// throws; in the original code, no target locals are written if control +// flow transfers out to an enclosing scope in the same function). +// +// Any operand that interferes must be written to a scratch local first and then +// copied to its target local after all operands have been evaluated. +class TupleInterferenceFinder { + std::vector interfering; + +public: + TupleInterferenceFinder(const ExpressionList& operands, + Index targetBase, + const PassOptions& passOptions, + const Module& wasm) + : interfering(operands.size(), false) { + Index numOperands = operands.size(); + assert(numOperands > 1); + + std::unordered_set subsequentReads; + std::unordered_set subsequentWrites; + bool subsequentTransfersControlFlow = false; + + for (Index i = numOperands; i > 0; i--) { + Index opIndex = i - 1; + Index targetLocal = targetBase + opIndex; + + if (subsequentTransfersControlFlow || + subsequentReads.contains(targetLocal) || + subsequentWrites.contains(targetLocal)) { + interfering[opIndex] = true; + } + + EffectAnalyzer effects(passOptions, wasm, operands[opIndex]); + subsequentReads.insert(effects.localsRead.begin(), + effects.localsRead.end()); + subsequentWrites.insert(effects.localsWritten.begin(), + effects.localsWritten.end()); + if (effects.transfersControlFlow()) { + subsequentTransfersControlFlow = true; + } + } + } + + bool interferes(Index i) const { return interfering[i]; } +}; + +} // anonymous namespace + struct TupleOptimization : public WalkerPass> { bool isFunctionParallel() override { return true; } @@ -231,15 +300,17 @@ struct TupleOptimization : public WalkerPass> { } } - MapApplier mapApplier(tupleToNewBaseMap); + MapApplier mapApplier(tupleToNewBaseMap, getPassOptions()); mapApplier.walkFunctionInModule(func, getModule()); } struct MapApplier : public PostWalker { std::unordered_map& tupleToNewBaseMap; + const PassOptions& passOptions; - MapApplier(std::unordered_map& tupleToNewBaseMap) - : tupleToNewBaseMap(tupleToNewBaseMap) {} + MapApplier(std::unordered_map& tupleToNewBaseMap, + const PassOptions& passOptions) + : tupleToNewBaseMap(tupleToNewBaseMap), passOptions(passOptions) {} // Gets the new base index if there is one, or 0 if not (0 is an impossible // value for a new index, as local index 0 was taken before, as tuple @@ -294,11 +365,31 @@ struct TupleOptimization : public WalkerPass> { auto* value = curr->value; if (auto* make = value->dynCast()) { - // Write each of the tuple.make fields into the proper local. + // If writing an operand directly to its target local would interfere + // with any subsequent operand (e.g. in a tuple swap), write it to a + // temporary local first and copy it at the end. + Index numOperands = type.size(); + TupleInterferenceFinder interferences( + make->operands, targetBase, passOptions, *getModule()); + + std::vector tempIndexes(numOperands); + for (Index i = 0; i < numOperands; i++) { + if (interferences.interferes(i)) { + tempIndexes[i] = Builder::addVar(getFunction(), type[i]); + } + } + std::vector sets; - for (Index i = 0; i < type.size(); i++) { - auto* value = make->operands[i]; - sets.push_back(builder.makeLocalSet(targetBase + i, value)); + for (Index i = 0; i < numOperands; i++) { + Index dest = + interferences.interferes(i) ? tempIndexes[i] : targetBase + i; + sets.push_back(builder.makeLocalSet(dest, make->operands[i])); + } + for (Index i = 0; i < numOperands; i++) { + if (interferences.interferes(i)) { + sets.push_back(builder.makeLocalSet( + targetBase + i, builder.makeLocalGet(tempIndexes[i], type[i]))); + } } replace(builder.makeBlock(sets)); return; diff --git a/test/lit/passes/tuple-optimization.wast b/test/lit/passes/tuple-optimization.wast index fce17ba0ba4..b9781159766 100644 --- a/test/lit/passes/tuple-optimization.wast +++ b/test/lit/passes/tuple-optimization.wast @@ -1065,7 +1065,7 @@ ) ) - ;; CHECK: (func $unreachable.tuple.extract (type $3) (result i32) + ;; CHECK: (func $unreachable.tuple.extract (type $4) (result i32) ;; CHECK-NEXT: (local $tuple (tuple i32 i64)) ;; CHECK-NEXT: (local $non-tuple i32) ;; CHECK-NEXT: (tuple.extract 2 0 @@ -1083,4 +1083,213 @@ ) ) ) + + ;; CHECK: (func $swap (type $3) (param $x i32) (param $y i32) (result i32) + ;; CHECK-NEXT: (local $t (tuple i32 i32)) + ;; CHECK-NEXT: (local $3 i32) + ;; CHECK-NEXT: (local $4 i32) + ;; CHECK-NEXT: (local $5 i32) + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $3 + ;; CHECK-NEXT: (local.get $x) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $4 + ;; CHECK-NEXT: (local.get $y) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $5 + ;; CHECK-NEXT: (local.get $4) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $4 + ;; CHECK-NEXT: (local.get $3) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $3 + ;; CHECK-NEXT: (local.get $5) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.get $4) + ;; CHECK-NEXT: ) + (func $swap (param $x i32) (param $y i32) (result i32) + ;; Swapping elements of a tuple must not overwrite the earlier element + ;; before the later element reads it. + (local $t (tuple i32 i32)) + (local.set $t + (tuple.make 2 + (local.get $x) + (local.get $y) + ) + ) + (local.set $t + (tuple.make 2 + (tuple.extract 2 1 + (local.get $t) + ) + (tuple.extract 2 0 + (local.get $t) + ) + ) + ) + (tuple.extract 2 1 + (local.get $t) + ) + ) + + ;; CHECK: (func $no-swap (type $3) (param $x i32) (param $y i32) (result i32) + ;; CHECK-NEXT: (local $t (tuple i32 i32)) + ;; CHECK-NEXT: (local $t' (tuple i32 i32)) + ;; CHECK-NEXT: (local $4 i32) + ;; CHECK-NEXT: (local $5 i32) + ;; CHECK-NEXT: (local $6 i32) + ;; CHECK-NEXT: (local $7 i32) + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $4 + ;; CHECK-NEXT: (local.get $x) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $5 + ;; CHECK-NEXT: (local.get $y) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $6 + ;; CHECK-NEXT: (local.get $5) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $7 + ;; CHECK-NEXT: (local.get $4) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.get $7) + ;; CHECK-NEXT: ) + (func $no-swap (param $x i32) (param $y i32) (result i32) + ;; Like $swap, but this time we are copying the elements to a different + ;; tuple local, so we don't need temp locals. + (local $t (tuple i32 i32)) + (local $t' (tuple i32 i32)) + (local.set $t + (tuple.make 2 + (local.get $x) + (local.get $y) + ) + ) + (local.set $t' + (tuple.make 2 + (tuple.extract 2 1 + (local.get $t) + ) + (tuple.extract 2 0 + (local.get $t) + ) + ) + ) + (tuple.extract 2 1 + (local.get $t') + ) + ) + + ;; CHECK: (func $swap-3 (type $5) (param $x i32) (param $y i32) (param $z i32) (result i32) + ;; CHECK-NEXT: (local $t (tuple i32 i32 i32)) + ;; CHECK-NEXT: (local $4 i32) + ;; CHECK-NEXT: (local $5 i32) + ;; CHECK-NEXT: (local $6 i32) + ;; CHECK-NEXT: (local $7 i32) + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $4 + ;; CHECK-NEXT: (local.get $x) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $5 + ;; CHECK-NEXT: (local.get $y) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $6 + ;; CHECK-NEXT: (local.get $z) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $7 + ;; CHECK-NEXT: (local.get $5) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $5 + ;; CHECK-NEXT: (local.get $6) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $6 + ;; CHECK-NEXT: (local.get $4) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $4 + ;; CHECK-NEXT: (local.get $7) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.get $6) + ;; CHECK-NEXT: ) + (func $swap-3 (param $x i32) (param $y i32) (param $z i32) (result i32) + ;; Rotating 3 elements of a tuple (0->1, 1->2, 2->0) requires saving only + ;; the element that would be overwritten before being read. + (local $t (tuple i32 i32 i32)) + (local.set $t + (tuple.make 3 + (local.get $x) + (local.get $y) + (local.get $z) + ) + ) + (local.set $t + (tuple.make 3 + (tuple.extract 3 1 + (local.get $t) + ) + (tuple.extract 3 2 + (local.get $t) + ) + (tuple.extract 3 0 + (local.get $t) + ) + ) + ) + (tuple.extract 3 2 + (local.get $t) + ) + ) + + ;; CHECK: (func $branch-out (type $6) (param $cond i32) (result i32) + ;; CHECK-NEXT: (local $t (tuple i32 i32)) + ;; CHECK-NEXT: (local $2 i32) + ;; CHECK-NEXT: (local $3 i32) + ;; CHECK-NEXT: (local $4 i32) + ;; CHECK-NEXT: (block $b + ;; CHECK-NEXT: (block + ;; CHECK-NEXT: (local.set $4 + ;; CHECK-NEXT: (i32.const 1) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $3 + ;; CHECK-NEXT: (block (result i32) + ;; CHECK-NEXT: (br_if $b + ;; CHECK-NEXT: (local.get $cond) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (i32.const 2) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.set $2 + ;; CHECK-NEXT: (local.get $4) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: ) + ;; CHECK-NEXT: (local.get $2) + ;; CHECK-NEXT: ) + (func $branch-out (param $cond i32) (result i32) + ;; If a later operand transfers control flow to an enclosing scope in the + ;; function, earlier operands must not prematurely overwrite target locals. + (local $t (tuple i32 i32)) + (block $b + (local.set $t + (tuple.make 2 + (i32.const 1) + (block (result i32) + (br_if $b (local.get $cond)) + (i32.const 2) + ) + ) + ) + ) + (tuple.extract 2 0 + (local.get $t) + ) + ) )