/usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/jit/tensorexpr
NameSizeModeActions
operators/-0755rm
analysis.h58880644editdlrm
block_codegen.h42110644editdlrm
bounds_inference.h22300644editdlrm
bounds_overlap.h33290644editdlrm
codegen.h64020644editdlrm
cpp_codegen.h22780644editdlrm
cpp_intrinsics.h7190644editdlrm
cuda_codegen.h77820644editdlrm
cuda_random.h26420644editdlrm
dim_arg.h8840644editdlrm
eval.h96390644editdlrm
exceptions.h32530644editdlrm
expr.h115880644editdlrm
external_functions.h12740644editdlrm
external_functions_registry.h23430644editdlrm
fwd_decls.h28060644editdlrm
graph_opt.h25530644editdlrm
half_support.h50380644editdlrm
hash_provider.h79300644editdlrm
intrinsic_symbols.h4200644editdlrm
ir.h226220644editdlrm
ir_cloner.h20690644editdlrm
ir_mutator.h20100644editdlrm
ir_printer.h36930644editdlrm
ir_simplifier.h150900644editdlrm
ir_verifier.h12400644editdlrm
ir_visitor.h18250644editdlrm
kernel.h92100644editdlrm
llvm_codegen.h31800644editdlrm
llvm_jit.h19650644editdlrm
loopnest.h215990644editdlrm
mem_dependency_checker.h130030644editdlrm
reduction.h67420644editdlrm
registerizer.h124980644editdlrm
stmt.h211380644editdlrm
tensor.h76400644editdlrm
tensorexpr_init.h2680644editdlrm
types.h38800644editdlrm
unique_name_manager.h9400644editdlrm
var_substitutor.h17530644editdlrm
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