/
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/half_support.h
(5038B)
#pragma once #include <torch/csrc/jit/tensorexpr/codegen.h> #include <torch/csrc/jit/tensorexpr/ir.h> #include <torch/csrc/jit/tensorexpr/ir_visitor.h> #include <torch/csrc/jit/tensorexpr/tensor.h> namespace torch { namespace jit { namespace tensorexpr { // Walk the Statment looking for Half size loads/stores. class HalfChecker : public IRVisitor { public: HalfChecker(const std::vector<CodeGen::BufferArg>& args) { for (const auto& BA : args) { hasHalf_ |= BA.dtype().scalar_type() == ScalarType::Half; } } bool hasHalf() const { return hasHalf_; } bool hasBFloat16() const { return hasBFloat16_; } void visit(LoadPtr v) override { hasHalf_ |= v->dtype().scalar_type() == ScalarType::Half; hasBFloat16_ |= v->dtype().scalar_type() == ScalarType::BFloat16; IRVisitor::visit(v); } void visit(StorePtr v) override { hasHalf_ |= v->buf()->dtype().scalar_type() == ScalarType::Half; hasBFloat16_ |= v->buf()->dtype().scalar_type() == ScalarType::BFloat16; IRVisitor::visit(v); } void visit(HalfImmPtr v) override { hasHalf_ = true; } void visit(BFloat16ImmPtr v) override { hasBFloat16_ = true; } void visit(CastPtr v) override { hasHalf_ |= v->dtype().scalar_type() == ScalarType::Half; hasBFloat16_ |= v->dtype().scalar_type() == ScalarType::BFloat16; IRVisitor::visit(v); } private: bool hasHalf_{false}; bool hasBFloat16_{false}; }; // NOLINTNEXTLINE(cppcoreguidelines-pro-type-member-init) class HalfRewriter : public IRMutator { ExprPtr mutate(LoadPtr v) override { ExprPtr child = IRMutator::mutate(v); if (!isHalf(child)) { return child; } ExprPtr ret = alloc<Cast>( child->dtype().cloneWithScalarType(ScalarType::Float), child); inserted_half_casts_.insert(ret); return ret; } StmtPtr mutate(StorePtr v) override { // Since mutation changes the `value()` expression in-place, we need to // get the dtype of the `value()` before that is mutated. auto newType = v->value()->dtype(); ExprPtr new_val = v->value()->accept_mutator(this); if (isHalf(newType.scalar_type())) { new_val = alloc<Cast>(newType, new_val); inserted_half_casts_.insert(new_val); } v->set_value(new_val); return v; } ExprPtr mutate(HalfImmPtr v) override { return alloc<Cast>(kFloat, v); } ExprPtr mutate(BFloat16ImmPtr v) override { return alloc<Cast>(kFloat, v); } ExprPtr mutate(CastPtr v) override { ExprPtr child = v->src_value()->accept_mutator(this); // just don't allow half casts we didn't insert. if (isHalf(v)) { if (inserted_half_casts_.count(v) < 1) { return child; } } // Remove Half(Float()) and friends. CastPtr cast_child = to<Cast>(child); if (cast_child) { if (v->dtype().is_floating_point() && cast_child->dtype().is_floating_point()) { return alloc<Cast>(v->dtype(), cast_child->src_value()); } } if (child == v->src_value()) { return v; } return alloc<Cast>(v->dtype(), child); } StmtPtr mutate(LetPtr v) override { if (isHalf(v->dtype().scalar_type())) { VarPtr load_new_var = alloc<Var>(v->var()->name_hint(), kFloat); ExprPtr new_value = alloc<Cast>( v->dtype().cloneWithScalarType(ScalarType::Float), v->value()->accept_mutator(this)); var_map[v->var()] = load_new_var; return alloc<Let>(load_new_var, new_value); } return IRMutator::mutate(v); } ExprPtr mutate(VarPtr v) override { auto it = var_map.find(v); if (it != var_map.end()) { return it->second; } return v; } template <typename T> ExprPtr mutateArithmetic(T v) { IRMutator::mutate(v); if (isHalf(v)) { v->set_dtype(v->dtype().cloneWithScalarType(c10::kFloat)); } return v; } ExprPtr mutate(AddPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(SubPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(MulPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(DivPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(MaxPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(MinPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(CompareSelectPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(BroadcastPtr v) override { return mutateArithmetic(v); } ExprPtr mutate(IfThenElsePtr v) override { return mutateArithmetic(v); } ExprPtr mutate(IntrinsicsPtr v) override { return mutateArithmetic(v); } private: static bool isHalf(ScalarType st) { return st == ScalarType::Half || st == ScalarType::BFloat16; } static bool isHalf(ExprPtr v) { return isHalf(v->dtype().scalar_type()); } std::unordered_set<ExprPtr> inserted_half_casts_; std::unordered_map<VarPtr, VarPtr> var_map; }; } // namespace tensorexpr } // namespace jit } // namespace torch
Save
cmd:
run