#include <algorithm>
#include <memory>
#include <ir/cost.h>
#include <ir/effects.h>
#include <ir/iteration.h>
#include <ir/linear-execution.h>
#include <ir/properties.h>
#include <ir/type-updating.h>
#include <ir/utils.h>
#include <pass.h>
#include <wasm-builder.h>
#include <wasm-traversal.h>
#include <wasm.h>
namespace wasm {
namespace {
struct HashedExpression {
Expression* expr;
size_t digest;
HashedExpression(Expression* expr, size_t digest)
: expr(expr), digest(digest) {}
HashedExpression(const HashedExpression& other)
: expr(other.expr), digest(other.digest) {}
};
struct HEHasher {
size_t operator()(const HashedExpression hashed) const {
return hashed.digest;
}
};
struct HEComparer {
bool operator()(const HashedExpression a, const HashedExpression b) const {
if (a.digest != b.digest) {
return false;
}
return ExpressionAnalyzer::equal(a.expr, b.expr);
}
};
using HashedExprs = std::unordered_map<HashedExpression,
SmallVector<Expression*, 1>,
HEHasher,
HEComparer>;
struct RequestInfo {
Index requests = 0;
Expression* original = nullptr;
void validate() const {
assert(!(requests && original));
assert(requests || original);
}
};
struct RequestInfoMap : public std::unordered_map<Expression*, RequestInfo> {
void dump(std::ostream& o) {
for (auto& [curr, info] : *this) {
o << *curr << " has " << info.requests << " reqs, orig: " << info.original
<< '\n';
}
}
};
struct Scanner
: public LinearExecutionWalker<Scanner, UnifiedExpressionVisitor<Scanner>> {
PassOptions& options;
RequestInfoMap& requestInfos;
Scanner(PassOptions& options, RequestInfoMap& requestInfos)
: options(options), requestInfos(requestInfos) {}
HashedExprs activeExprs;
SmallVector<size_t, 10> activeHashes;
static void doNoteNonLinear(Scanner* self, Expression** currp) {
self->activeExprs.clear();
self->activeHashes.clear();
}
void visitExpression(Expression* curr) {
auto numChildren = Properties::getNumChildren(curr);
auto hash = ExpressionAnalyzer::shallowHash(curr);
for (Index i = 0; i < numChildren; i++) {
if (activeHashes.empty()) {
return;
}
hash_combine(hash, activeHashes.back());
activeHashes.pop_back();
}
activeHashes.push_back(hash);
if (!isRelevant(curr)) {
return;
}
auto& vec = activeExprs[HashedExpression(curr, hash)];
vec.push_back(curr);
if (vec.size() > 1) {
auto& info = requestInfos[curr];
auto* original = vec[0];
info.original = original;
requestInfos[original].requests++;
for (auto* child : ChildIterator(curr)) {
if (!requestInfos.count(child)) {
continue;
}
auto& childInfo = requestInfos[child];
auto* childOriginal = childInfo.original;
requestInfos.erase(child);
assert(childOriginal);
auto& childOriginalRequests = requestInfos[childOriginal].requests;
assert(childOriginalRequests > 0);
childOriginalRequests--;
if (childOriginalRequests == 0) {
requestInfos.erase(childOriginal);
}
}
}
}
bool isRelevant(Expression* curr) {
if (!curr->type.isConcrete() || curr->is<LocalGet>() ||
curr->is<LocalSet>() || Properties::isConstantExpression(curr) ||
!TypeUpdating::canHandleAsLocal(curr->type)) {
return false;
}
if (options.shrinkLevel > 0 && Measurer::measure(curr) >= 3) {
return true;
}
if (options.shrinkLevel == 0 && CostAnalyzer(curr).cost > 0) {
return true;
}
return false;
}
};
struct Checker
: public LinearExecutionWalker<Checker, UnifiedExpressionVisitor<Checker>> {
PassOptions& options;
RequestInfoMap& requestInfos;
Checker(PassOptions& options, RequestInfoMap& requestInfos)
: options(options), requestInfos(requestInfos) {}
struct ActiveOriginalInfo {
Index requestsLeft;
EffectAnalyzer effects;
};
std::unordered_map<Expression*, ActiveOriginalInfo> activeOriginals;
void visitExpression(Expression* curr) {
assert(!activeOriginals.count(curr));
if (!activeOriginals.empty()) {
EffectAnalyzer effects(options, *getModule());
effects.visit(curr);
std::vector<Expression*> invalidated;
for (auto& kv : activeOriginals) {
auto* original = kv.first;
auto& originalInfo = kv.second;
if (effects.invalidates(originalInfo.effects)) {
invalidated.push_back(original);
}
}
for (auto* original : invalidated) {
requestInfos[original].requests -=
activeOriginals.at(original).requestsLeft;
if (requestInfos[original].requests == 0) {
requestInfos.erase(original);
}
activeOriginals.erase(original);
}
}
auto iter = requestInfos.find(curr);
if (iter == requestInfos.end()) {
return;
}
auto& info = iter->second;
info.validate();
if (info.requests > 0) {
EffectAnalyzer effects(options, *getModule(), curr);
effects.trap = false;
if (effects.hasSideEffects() ||
Properties::isGenerative(curr, getModule()->features)) {
requestInfos.erase(curr);
} else {
activeOriginals.emplace(
curr, ActiveOriginalInfo{info.requests, std::move(effects)});
}
} else if (info.original) {
auto originalIter = activeOriginals.find(info.original);
if (originalIter == activeOriginals.end()) {
requestInfos.erase(iter);
return;
}
auto& originalInfo = originalIter->second;
if (originalInfo.requestsLeft == 1) {
activeOriginals.erase(info.original);
} else {
originalInfo.requestsLeft--;
}
}
}
static void doNoteNonLinear(Checker* self, Expression** currp) {
assert(self->activeOriginals.empty());
}
void visitFunction(Function* curr) {
assert(activeOriginals.empty());
}
};
struct Applier
: public LinearExecutionWalker<Applier, UnifiedExpressionVisitor<Applier>> {
RequestInfoMap requestInfos;
Applier(RequestInfoMap& requestInfos) : requestInfos(requestInfos) {}
std::unordered_map<Expression*, Index> originalLocalMap;
void visitExpression(Expression* curr) {
auto iter = requestInfos.find(curr);
if (iter == requestInfos.end()) {
return;
}
const auto& info = iter->second;
info.validate();
if (info.requests) {
Index local = originalLocalMap[curr] =
Builder::addVar(getFunction(), curr->type);
replaceCurrent(
Builder(*getModule()).makeLocalTee(local, curr, curr->type));
} else if (info.original) {
auto& originalInfo = requestInfos.at(info.original);
if (originalInfo.requests) {
assert(originalLocalMap.count(info.original));
replaceCurrent(
Builder(*getModule())
.makeLocalGet(originalLocalMap[info.original], curr->type));
originalInfo.requests--;
}
}
}
static void doNoteNonLinear(Applier* self, Expression** currp) {
self->originalLocalMap.clear();
}
};
}
struct LocalCSE : public WalkerPass<PostWalker<LocalCSE>> {
bool isFunctionParallel() override { return true; }
bool invalidatesDWARF() override { return true; }
std::unique_ptr<Pass> create() override {
return std::make_unique<LocalCSE>();
}
void doWalkFunction(Function* func) {
auto& options = getPassOptions();
RequestInfoMap requestInfos;
Scanner scanner(options, requestInfos);
scanner.walkFunctionInModule(func, getModule());
if (requestInfos.empty()) {
return;
}
Checker checker(options, requestInfos);
checker.walkFunctionInModule(func, getModule());
if (requestInfos.empty()) {
return;
}
Applier applier(requestInfos);
applier.walkFunctionInModule(func, getModule());
}
};
Pass* createLocalCSEPass() { return new LocalCSE(); }
}