/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/block_codegen.h (4211B)
#pragma once #include #include #include #include #include #include #include #include #include #include #include namespace torch { namespace jit { namespace tensorexpr { // A class that analyzes the given program relevant for Block backend. class BlockAnalysis : public IRVisitor { public: bool is_buf_store_target(BufPtr buf) const { return store_targets_.count(buf) > 0; } const std::unordered_set& loads() const { return loads_; } const std::unordered_set& stores() const { return store_targets_; } int block_size() const { return block_size_; } bool areBufsInMap(const std::unordered_set& bufs) const; BufPtr getMultiDimBuf(BufPtr buf) const; std::string getInputName(BufPtr buf) const; std::string getFlatInputName(BufPtr buf) const { return getInputName(buf) + "_flat"; } std::unordered_map getBufferMap() const { return map_input_to_tensor_bufs_; } private: void visit(StorePtr v) override; void visit(LoadPtr v) override; void visit(ForPtr v) override; std::unordered_map map_input_to_tensor_bufs_; std::unordered_set store_targets_; std::unordered_set loads_; int block_size_ = 32; }; // A class that overrides the underlying IRPrinter to produce Block. class BlockPrinter : public IRPrinter { public: BlockPrinter(std::ostream* os, BlockAnalysis* block_analysis) : IRPrinter(*os), block_analysis_(block_analysis) {} using IRPrinter::name_manager; using IRPrinter::visit; private: BlockAnalysis* block_analysis_; std::unordered_map dim_values_map; std::vector dim_names = {"N", "H", "W", "C"}; std::vector flat_dim_names = {"N", "NH", "NHW", "NHWC"}; void PrintTensorInfo(const std::unordered_set& bufs); void PrintArguments(const std::unordered_set& bufs); void PrintBufferInfo(const std::unordered_set& bufs); void PrintDistribution(const std::unordered_set& bufs); void PrintLoop(const std::unordered_set& bufs, bool block_idx = true); void PrintReshapeInfo( const std::unordered_set& bufs, bool reverse = false); void PrintDMAs(const std::unordered_set& bufs); void PrintAdjustBuffers(const std::unordered_set& bufs); void visit(ForPtr v) override; void visit(LoadPtr v) override; void visit(StorePtr v) override; void visit(BlockPtr v) override; void visit(AddPtr v) override; void visit(MulPtr v) override; }; class TORCH_API BlockCodeGen : public CodeGen { public: template /* implicit */ BlockCodeGen(StmtPtr stmt, Ts... ts) : CodeGen( stmt, std::vector({BufferArg(ts)...}), at::Device(at::kCPU)) { Initialize(); } BlockCodeGen( StmtPtr stmt, const std::vector& buffer_args, at::Device device = at::Device(at::kCPU), const std::string& kernel_func_name = "func") : CodeGen(stmt, buffer_args, device, kernel_func_name) { Initialize(); } ~BlockCodeGen() override; void call(const std::vector& args) override; void call_raw(const std::vector& args) override; void Initialize(); std::string getCodeText(const std::string& attr = "") override { return oss_.str(); } private: UniqueNameManager* name_manager() { if (!printer_) { throw std::runtime_error("Null IRPrinter is not expected"); } return printer_->name_manager(); } std::ostream& os() { return printer_->os(); } std::ostringstream oss_; std::unique_ptr printer_; std::unique_ptr block_analysis_; std::string GetUniqueFuncName(const std::string& func_prefix); }; } // namespace tensorexpr } // namespace jit } // namespace torch