/
usr
/
local
/
lib64
/
python3.6
/
site-packages
/
torch
/
include
/
torch
/
csrc
/
jit
/
tensorexpr
/
/usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/jit/tensorexpr
mkdir
upload
Name
Size
Mode
Actions
operators/
-
0755
rm
analysis.h
5888
0644
edit
dl
rm
block_codegen.h
4211
0644
edit
dl
rm
bounds_inference.h
2230
0644
edit
dl
rm
bounds_overlap.h
3329
0644
edit
dl
rm
codegen.h
6402
0644
edit
dl
rm
cpp_codegen.h
2278
0644
edit
dl
rm
cpp_intrinsics.h
719
0644
edit
dl
rm
cuda_codegen.h
7782
0644
edit
dl
rm
cuda_random.h
2642
0644
edit
dl
rm
dim_arg.h
884
0644
edit
dl
rm
eval.h
9639
0644
edit
dl
rm
exceptions.h
3253
0644
edit
dl
rm
expr.h
11588
0644
edit
dl
rm
external_functions.h
1274
0644
edit
dl
rm
external_functions_registry.h
2343
0644
edit
dl
rm
fwd_decls.h
2806
0644
edit
dl
rm
graph_opt.h
2553
0644
edit
dl
rm
half_support.h
5038
0644
edit
dl
rm
hash_provider.h
7930
0644
edit
dl
rm
intrinsic_symbols.h
420
0644
edit
dl
rm
ir.h
22622
0644
edit
dl
rm
ir_cloner.h
2069
0644
edit
dl
rm
ir_mutator.h
2010
0644
edit
dl
rm
ir_printer.h
3693
0644
edit
dl
rm
ir_simplifier.h
15090
0644
edit
dl
rm
ir_verifier.h
1240
0644
edit
dl
rm
ir_visitor.h
1825
0644
edit
dl
rm
kernel.h
9210
0644
edit
dl
rm
llvm_codegen.h
3180
0644
edit
dl
rm
llvm_jit.h
1965
0644
edit
dl
rm
loopnest.h
21599
0644
edit
dl
rm
mem_dependency_checker.h
13003
0644
edit
dl
rm
reduction.h
6742
0644
edit
dl
rm
registerizer.h
12498
0644
edit
dl
rm
stmt.h
21138
0644
edit
dl
rm
tensor.h
7640
0644
edit
dl
rm
tensorexpr_init.h
268
0644
edit
dl
rm
types.h
3880
0644
edit
dl
rm
unique_name_manager.h
940
0644
edit
dl
rm
var_substitutor.h
1753
0644
edit
dl
rm
Edit:
/usr/local/lib64/python3.6/site-packages/torch/include/torch/csrc/jit/tensorexpr/cpp_codegen.h
(2278B)
#pragma once #include <torch/csrc/jit/tensorexpr/codegen.h> #include <torch/csrc/jit/tensorexpr/ir_printer.h> namespace torch { namespace jit { namespace tensorexpr { class CppVarNameRewriter; // Generates C++ code from the IR. // // Vector operations are unrolled. // For example: // C[Ramp(0, 1, 3)] = A[Ramp(0, 2, 3)] + B[Ramp(0, 3, 3)]; // is unrolled into: // C[0] = A[0] + B[0]; // C[1] = A[2] + B[3]; // C[2] = A[4] + B[6]; class TORCH_API CppPrinter : public IRPrinter { public: explicit CppPrinter(std::ostream* os); ~CppPrinter() override; void printPrologue(); using IRPrinter::visit; // Binary expressions. void visit(ModPtr) override; void visit(MaxPtr) override; void visit(MinPtr) override; // Conditional expressions. void visit(CompareSelectPtr) override; void visit(IfThenElsePtr) override; // Tensor operations. void visit(AllocatePtr) override; void visit(FreePtr) override; void visit(LoadPtr) override; void visit(StorePtr) override; // Casts. void visit(CastPtr) override; void visit(BitCastPtr) override; // Calls. void visit(IntrinsicsPtr) override; void visit(ExternalCallPtr) override; // Vars. void visit(LetPtr) override; void visit(VarPtr) override; // Vector data types. void visit(RampPtr) override; void visit(BroadcastPtr) override; private: int lane_; std::unordered_map<VarPtr, ExprPtr> vector_vars_; }; class TORCH_API CppCodeGen : public CodeGen { public: CppCodeGen( StmtPtr stmt, const std::vector<BufferArg>& buffer_args, at::Device device = at::kCPU, const std::string& kernel_func_name = "func"); ~CppCodeGen() override; void call(const std::vector<CallArg>& args) override; void call_raw(const std::vector<void*>& args) override; template <typename... Ts> void operator()(const Ts&... ts) { call(std::vector<CallArg>({CallArg(ts)...})); } std::string getCodeText(const std::string& attr = "") override { return oss_.str(); } private: void init(); std::ostream& os() { return printer_->os(); } std::ostringstream oss_; std::unique_ptr<CppPrinter> printer_; std::unique_ptr<CppVarNameRewriter> var_name_rewriter_; }; } // namespace tensorexpr } // namespace jit } // namespace torch
Save
cmd:
run