/usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/jit/tensorexpr
Edit: /usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/jit/tensorexpr/analysis.h (5888B)
#pragma once
#include
#include
#include
#include
namespace torch {
namespace jit {
namespace tensorexpr {
class HasRand : public IRVisitor {
public:
HasRand(StmtPtr stmt) : stmt_(stmt) {
stmt_->accept(this);
}
bool has_rand() const {
return has_rand_;
}
private:
void visit(IntrinsicsPtr v) override {
if (v->op_type() == IntrinsicsOp::kRand) {
has_rand_ = true;
} else {
IRVisitor::visit(v);
}
}
StmtPtr stmt_;
bool has_rand_ = false;
};
template
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
class NodeFinder : public IRVisitor {
public:
void visit(NodePtr v) override {
nodes.push_back((NodePtr)v);
IRVisitor::visit(v);
}
static std::vector> find(StmtPtr s) {
NodeFinder nf;
s->accept(&nf);
return nf.nodes;
}
static std::vector> find(ExprPtr e) {
NodeFinder nf;
e->accept(&nf);
return nf.nodes;
}
std::vector> nodes;
};
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
class VarFinder : public IRVisitor {
public:
void visit(VarPtr v) override {
vars_.insert(v);
IRVisitor::visit(v);
}
static std::unordered_set find(StmtPtr s) {
VarFinder nf;
s->accept(&nf);
return nf.vars();
}
static std::unordered_set find(ExprPtr e) {
VarFinder nf;
e->accept(&nf);
return nf.vars();
}
const std::unordered_set& vars() {
return vars_;
}
private:
std::unordered_set vars_;
};
class BufFinder : public IRVisitor {
public:
void visit(BufPtr v) override {
bufs_.insert(v);
IRVisitor::visit(v);
}
static std::unordered_set find(StmtPtr s) {
BufFinder nf;
s->accept(&nf);
return nf.bufs();
}
static std::unordered_set find(ExprPtr e) {
BufFinder nf;
e->accept(&nf);
return nf.bufs();
}
const std::unordered_set& bufs() {
return bufs_;
}
private:
std::unordered_set bufs_;
};
// Finds all kinds of write operations to the provided Buf.
class WritesToBuf : public IRVisitor {
public:
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
WritesToBuf(BufPtr target) : target_(target) {}
std::vector writes() {
return writes_;
}
static std::vector find(StmtPtr s, BufPtr b) {
WritesToBuf finder(b);
s->accept(&finder);
return finder.writes();
}
private:
void visit(StorePtr v) override {
if (v->buf() == target_) {
writes_.push_back(v);
}
}
void visit(AtomicAddPtr v) override {
if (v->buf() == target_) {
writes_.push_back(v);
}
}
BufPtr target_;
std::vector writes_;
};
class StmtsReadingBuf : public IRVisitor {
public:
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
StmtsReadingBuf(BufPtr target) : target_(target) {}
std::vector reads() {
return reads_;
}
static std::vector find(StmtPtr s, BufPtr b) {
StmtsReadingBuf finder(b);
s->accept(&finder);
return finder.reads();
}
private:
bool readsBuffer(StmtPtr s) {
auto loads = NodeFinder::find(s);
for (auto l : loads) {
if (l->buf() == target_) {
return true;
}
}
return false;
}
void visit(StorePtr v) override {
if (readsBuffer(v)) {
reads_.push_back(v);
}
}
void visit(LetPtr v) override {
if (readsBuffer(v)) {
reads_.push_back(v);
}
}
void visit(CondPtr v) override {
if (readsBuffer(v)) {
reads_.push_back(v);
}
}
void visit(AtomicAddPtr v) override {
if (readsBuffer(v)) {
reads_.push_back(v);
}
}
BufPtr target_;
std::vector reads_;
};
// Traverses the IR to determine if a particular Var is modified within it.
class ModifiesVarChecker : public IRVisitor {
public:
ModifiesVarChecker(VarPtr v) : var_(v) {}
static bool check(StmtPtr s, VarPtr v) {
ModifiesVarChecker checker(v);
s->accept(&checker);
return checker.found();
}
bool found() {
return found_;
}
private:
void visit(StorePtr v) override {
if (v->buf()->base_handle() == var_) {
found_ = true;
return;
}
IRVisitor::visit(v);
}
void visit(AtomicAddPtr v) override {
if (v->buf()->base_handle() == var_) {
found_ = true;
return;
}
IRVisitor::visit(v);
}
void visit(LetPtr v) override {
if (v->var() == var_) {
found_ = true;
return;
}
IRVisitor::visit(v);
}
void visit(ForPtr v) override {
if (v->var() == var_) {
found_ = true;
return;
}
IRVisitor::visit(v);
}
VarPtr var_;
bool found_{false};
};
// A class that analyzes the given program relevant for Block backend
// It creates a map of multi dim buffers and their flat verions
class CreateBufferMap : public IRVisitor {
public:
const std::unordered_map& getBufferMap() const {
return map_input_to_tensor_bufs_;
}
private:
void visit(StorePtr v) override {
auto load_node = to(v->value());
if (load_node) {
auto t_buf = load_node->buf();
map_input_to_tensor_bufs_.emplace(t_buf->name_hint(), v->buf());
} else {
auto add_node = to(v->value());
auto mul_node = to(v->value());
// This means for now, v->value() can be Add or Mul
TORCH_INTERNAL_ASSERT(add_node || mul_node, buildErrorMessage());
map_input_to_tensor_bufs_.emplace(v->buf()->name_hint(), v->buf());
}
v->value()->accept(this);
}
std::unordered_map map_input_to_tensor_bufs_;
};
} // namespace tensorexpr
} // namespace jit
} // namespace torch