/usr/local/lib64/python3.6/site-packages/torch/include/caffe2/core
NameSizeModeActions
allocator.h1360644editdlrm
blob.h41680644editdlrm
blob_serialization.h107910644editdlrm
blob_serializer_base.h39050644editdlrm
blob_stats.h11270644editdlrm
common.h43290644editdlrm
common_cudnn.h98930644editdlrm
common_gpu.h214140644editdlrm
common_omp.h1560644editdlrm
context.h61740644editdlrm
context_base.h43820644editdlrm
context_gpu.h110140644editdlrm
cudnn_wrappers.h69560644editdlrm
db.h93520644editdlrm
distributions_stubs.h21610644editdlrm
event.h124200644editdlrm
event_cpu.h11920644editdlrm
export_c10_op_to_caffe2.h94870644editdlrm
export_caffe2_op_to_c10.h111010644editdlrm
flags.h740644editdlrm
graph.h52580644editdlrm
init.h64960644editdlrm
logging.h750644editdlrm
macros.h34260644editdlrm
memonger.h8170644editdlrm
module.h24730644editdlrm
net.h46340644editdlrm
net_async_base.h73970644editdlrm
net_async_scheduling.h9930644editdlrm
net_async_task.h8330644editdlrm
net_async_task_future.h19250644editdlrm
net_async_task_graph.h22530644editdlrm
net_async_tracing.h50930644editdlrm
net_dag_utils.h21460644editdlrm
net_parallel.h21440644editdlrm
net_simple.h26060644editdlrm
net_simple_refcount.h20970644editdlrm
numa.h720644editdlrm
observer.h38090644editdlrm
operator.h588720644editdlrm
operator_gradient.h102220644editdlrm
operator_schema.h184770644editdlrm
plan_executor.h2190644editdlrm
prof_dag_counters.h27510644editdlrm
qtensor.h66150644editdlrm
qtensor_serialization.h26240644editdlrm
scope_guard.h46750644editdlrm
static_tracepoint.h3980644editdlrm
static_tracepoint_elfx86.h55550644editdlrm
stats.h103650644editdlrm
storage.h7330644editdlrm
tensor.h186680644editdlrm
tensor_impl.h3510644editdlrm
tensor_int8.h4500644editdlrm
test_utils.h62850644editdlrm
timer.h12180644editdlrm
transform.h57410644editdlrm
types.h22480644editdlrm
workspace.h113050644editdlrm
Edit: /usr/local/lib64/python3.6/site-packages/torch/include/caffe2/core/net.h (4634B)
#ifndef CAFFE2_CORE_NET_H_ #define CAFFE2_CORE_NET_H_ #include #include #include #include // NOLINT #include #include #include #include "c10/core/thread_pool.h" #include "c10/util/Registry.h" #include "caffe2/core/blob.h" #include "caffe2/core/common.h" #include "caffe2/core/logging.h" #include "caffe2/core/observer.h" #include "caffe2/core/operator_schema.h" #include "caffe2/core/tensor.h" #include "caffe2/proto/caffe2_pb.h" #include "caffe2/utils/simple_queue.h" C10_DECLARE_string(caffe2_override_executor); namespace caffe2 { class NetBase; typedef ObserverBase NetObserver; typedef std::function(NetBase*)> NetObserverCreator; class OperatorBase; class Workspace; // Net is a thin struct that owns all the operators together with the operator // contexts. class TORCH_API NetBase : public Observable { public: NetBase(const std::shared_ptr& net_def, Workspace* ws); virtual ~NetBase() noexcept {} virtual bool SupportsAsync() = 0; inline const vector& events() const { return events_; } virtual void Wait() { // by default just wait till all events are finished for (const auto& event : events_) { event->Finish(); } } virtual bool Run() { if (!RunAsync()) { LOG(ERROR) << "Failed to execute async run"; return false; } Wait(); return handleRunError(); } virtual bool RunAsync(); virtual void Cancel(); /* Benchmarks a network for one individual run so that we can feed new * inputs on additional calls. * This function returns the number of microseconds spent * during the benchmark */ virtual float TEST_Benchmark_One_Run(); /** * Benchmarks a network. * * This function returns a vector of float recording the number of milli- * seconds spent during the benchmark. The 0-th item is the time spent per * each network run, and if a net instantiation supports run_individual, * the remainder of the vector returns the number of milliseconds spent per * operator. */ virtual vector TEST_Benchmark( const int /*warmup_runs*/, const int /*main_runs*/, const bool /*run_individual*/); inline const vector& external_output() const { return external_output_; } inline const vector& external_input() const { return external_input_; } /* Used to attach Observers to operators of a Net * * Returns pointers to objects owned with unique_ptrs. * Use with caution. */ virtual vector GetOperators() const = 0; const string& Name() const { return name_; } inline const NetDef& debug_def() const { CAFFE_ENFORCE(has_debug_def(), "net_def was null!"); return *net_def_; } inline bool has_debug_def() const { return net_def_ != nullptr; } protected: virtual bool DoRunAsync() { CAFFE_THROW("Not implemented"); }; virtual bool handleRunError() { for (const Event* event : events_) { if (event->Query() != EventStatus::EVENT_SUCCESS) { CAFFE_THROW(event->ErrorMessage()); } } return true; } vector external_input_; vector external_output_; string name_; vector events_; std::shared_ptr net_def_; C10_DISABLE_COPY_AND_ASSIGN(NetBase); }; class TORCH_API ExecutorHelper { public: ExecutorHelper() {} virtual TaskThreadPoolBase* GetPool(const DeviceOption& option) const; virtual std::vector GetOperators() const; virtual int GetNumWorkers() const; virtual ~ExecutorHelper() {} }; C10_DECLARE_REGISTRY( NetRegistry, NetBase, const std::shared_ptr&, Workspace*); #define REGISTER_NET_CREATOR(key, ...) \ C10_REGISTER_CREATOR(NetRegistry, key, __VA_ARGS__) #define REGISTER_NET(name, ...) \ C10_REGISTER_CLASS(NetRegistry, name, __VA_ARGS__) /** * @brief Creates a network, accessing / creating blobs in the given workspace. * * Note that this is different from Workspace::CreateNet. The latter adds the * created net object to the workspace's net map, while this function returns * a standalone net object. */ TORCH_API unique_ptr CreateNet(const NetDef& net_def, Workspace* ws); TORCH_API unique_ptr CreateNet( const std::shared_ptr& net_def, Workspace* ws); TORCH_API void AddGlobalNetObserverCreator(NetObserverCreator creator); TORCH_API void ClearGlobalNetObservers(); } // namespace caffe2 #endif // CAFFE2_CORE_NET_H_