/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/codegen.h (6402B)
#pragma once
#include
#include
#include
namespace torch {
namespace jit {
namespace tensorexpr {
template
class PaddedBuffer;
class TORCH_API CodeGen {
public:
class BufferArg;
class CallArg;
template
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
CodeGen(StmtPtr stmt, Ts... ts)
: stmt_(stmt), buffer_args_({BufferArg(ts)...}) {}
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
CodeGen(
StmtPtr stmt,
std::vector buffer_args,
at::Device device = at::kCPU,
std::string kernel_func_name = "func")
: stmt_(stmt),
buffer_args_(std::move(buffer_args)),
device_(device),
kernel_func_name_(std::move(kernel_func_name)) {}
virtual ~CodeGen() = default;
StmtPtr stmt() const {
return stmt_;
}
void set_stmt(StmtPtr s) {
stmt_ = s;
}
void apply_mutator(IRMutator* mutator) {
stmt_ = stmt_->accept_mutator(mutator);
}
void apply_visitor(IRVisitor* visitor) {
stmt_->accept(visitor);
}
std::vector& buffer_args() {
return buffer_args_;
}
const std::vector& buffer_args() const {
return buffer_args_;
}
at::Device device() {
return device_;
}
// This function returns the generated code as
// a string.
virtual std::string getCodeText(const std::string& attr = "") {
return ("");
}
// There are two ways to invoke the codegen:
// 1) with a vector of CallArgs
// 2) with a vector of raw 'void*' pointers
//
// The codegen knows types of all inputs from the buffer args, that's why
// 'void*' pointers suffice.
//
// TODO: Eventually we might consider killing the CallArgs version, but
// currently only LLVM codegen implements call_raw.
virtual void call(const std::vector& args) = 0;
virtual void call_raw(const std::vector& args) = 0;
virtual at::Tensor empty_strided(
c10::IntArrayRef size,
c10::IntArrayRef stride,
c10::optional dtype_opt,
c10::optional layout_opt,
c10::optional device_opt,
c10::optional pin_memory_opt) {
return at::empty_strided(
size, stride, dtype_opt, layout_opt, device_opt, pin_memory_opt);
}
const std::string& kernel_func_name() const {
return kernel_func_name_;
}
protected:
static void* argToPtr(const BufferArg& bufferArg, const CallArg& callArg);
private:
StmtPtr stmt_;
std::vector buffer_args_;
at::Device device_ = at::kCPU;
std::string kernel_func_name_ = "func";
};
class CodeGen::BufferArg {
public:
BufferArg(Tensor tensor) : buf_(tensor.buf()) {}
BufferArg(const VarHandle& var) : var_(var.node()), isVar_(true) {}
BufferArg(const BufHandle& buf) : buf_(buf.node()) {}
VarPtr var() const {
return isVar_ ? var_ : buf_->base_handle();
}
BufPtr buf() const {
return buf_;
}
bool isVar() const {
return isVar_;
}
Dtype dtype() const {
return isVar_ ? var_->dtype() : buf_->dtype();
}
private:
VarPtr var_ = nullptr;
BufPtr buf_ = nullptr;
bool isVar_ = false;
};
class CodeGen::CallArg {
public:
template
CallArg(const PaddedBuffer& buffer);
template
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init,cppcoreguidelines-pro-type-const-cast)
CallArg(const std::vector& buffer)
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-const-cast)
: data_(const_cast(buffer.data())) {}
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
CallArg(void* ptr) : data_(ptr) {}
#define ARG_TYPE_CTOR(Type, Name) \
CallArg(Type v) { \
memcpy(&data_, &v, sizeof(Type)); \
}
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
AT_FORALL_SCALAR_TYPES_AND3(Bool, Half, BFloat16, ARG_TYPE_CTOR);
#undef ARG_TYPE_CTOR
void* data() const {
return data_;
}
#define ARG_PTR_DEFINE(Type, Name) \
Type* Name##Ptr() const { \
return (Type*)&data_; \
}
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-const-cast)
AT_FORALL_SCALAR_TYPES_AND3(Bool, Half, BFloat16, ARG_PTR_DEFINE);
#undef ARG_PTR_DEFINE
private:
void* data_;
};
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
class RegisterCodeGenList {
public:
TORCH_API static RegisterCodeGenList& GetInstance() {
static RegisterCodeGenList codegen_list;
return codegen_list;
}
using StmtFactoryMethod = std::function(
StmtPtr stmt,
const std::vector&,
at::Device device,
const std::string& kernel_func_name)>;
TORCH_API StmtFactoryMethod FindStmtFactoryMethod(const std::string& name);
RegisterCodeGenList(const RegisterCodeGenList&) = delete;
RegisterCodeGenList& operator=(const RegisterCodeGenList&) = delete;
private:
template
friend class RegisterCodeGen;
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init)
RegisterCodeGenList() = default;
TORCH_API void AddStmtFactoryMethod(
const std::string& name,
const StmtFactoryMethod& stmt_factory_method);
std::unordered_map stmt_factory_methods_;
};
template
class RegisterCodeGen {
public:
explicit RegisterCodeGen(const std::string& name) {
RegisterCodeGenList& codegen_list = RegisterCodeGenList::GetInstance();
codegen_list.AddStmtFactoryMethod(
name,
[](StmtPtr stmt,
const std::vector& params,
at::Device device,
const std::string& kernel_func_name) {
// NOLINTNEXTLINE(cppcoreguidelines-init-variables)
std::unique_ptr method(
new CodeGenType(stmt, params, device, kernel_func_name));
return method;
});
}
};
TORCH_API std::unique_ptr CreateCodeGen(
const std::string& name,
StmtPtr stmt,
const std::vector& params,
at::Device device = at::kCPU,
const std::string& kernel_func_name = "func");
class TORCH_API GenericIntrinsicsExpander : public IRMutator {
protected:
ExprPtr mutate(IntrinsicsPtr v) override;
};
} // namespace tensorexpr
} // namespace jit
} // namespace torch