/usr/local/lib64/python3.6/site-packages/torch/include/ATen
NameSizeModeActions
core/-0755rm
cpu/-0755rm
cuda/-0755rm
cudnn/-0755rm
detail/-0755rm
hip/-0755rm
native/-0755rm
quantized/-0755rm
AccumulateType.h44380644editdlrm
ArrayRef.h440644editdlrm
ATen.h9980644editdlrm
autocast_mode.h67160644editdlrm
Backend.h430644editdlrm
Backtrace.h460644editdlrm
BatchedFallback.h9650644editdlrm
BatchedTensorImpl.h53830644editdlrm
CompositeExplicitAutogradFunctions.h16220644editdlrm
CompositeExplicitAutogradFunctions_inl.h540750644editdlrm
CompositeImplicitAutogradFunctions.h16220644editdlrm
CompositeImplicitAutogradFunctions_inl.h1420820644editdlrm
Config.h7340644editdlrm
Context.h127670644editdlrm
cpp_custom_type_hack.h53260644editdlrm
CPUApplyUtils.h125820644editdlrm
CPUFixedAllocator.h8300644editdlrm
CPUFunctions.h16000644editdlrm
CPUFunctions_inl.h1719240644editdlrm
CPUGeneratorImpl.h14310644editdlrm
CUDAFunctions.h16010644editdlrm
CUDAFunctions_inl.h1856960644editdlrm
CUDAGeneratorImpl.h46950644editdlrm
Device.h420644editdlrm
DeviceGuard.h11340644editdlrm
Dimname.h310644editdlrm
DimVector.h460644editdlrm
Dispatch.h521370644editdlrm
div_rtn.h2040644editdlrm
DLConvertor.h5760644editdlrm
dlpack.h52440644editdlrm
DynamicLibrary.h3690644editdlrm
ExpandUtils.h145060644editdlrm
Formatting.h340644editdlrm
Functions.h8463260644editdlrm
Generator.h460644editdlrm
InferSize.h21430644editdlrm
InitialTensorOptions.h4450644editdlrm
Layout.h420644editdlrm
MapAllocator.h29990644editdlrm
MatrixRef.h30160644editdlrm
MemoryOverlap.h11170644editdlrm
MetaFunctions.h16010644editdlrm
MetaFunctions_inl.h840060644editdlrm
NamedTensor.h350644editdlrm
NamedTensorUtils.h57470644editdlrm
NativeFunctions.h3546510644editdlrm
NativeMetaFunctions.h354450644editdlrm
NumericUtils.h27870644editdlrm
OpaqueTensorImpl.h60800644editdlrm
Operators.h17071990644editdlrm
OpMathType.h4600644editdlrm
Parallel.h48750644editdlrm
ParallelNative.h24430644editdlrm
ParallelNativeTBB.h29340644editdlrm
ParallelOpenMP.h30490644editdlrm
PTThreadPool.h3940644editdlrm
record_function.h240440644editdlrm
RedispatchFunctions.h11128860644editdlrm
RegistrationDeclarations.h5457770644editdlrm
SavedTensorHooks.h3280644editdlrm
Scalar.h440644editdlrm
ScalarOps.h22720644editdlrm
ScalarType.h1290644editdlrm
SequenceNumber.h3730644editdlrm
SmallVector.h470644editdlrm
SparseCsrTensorImpl.h20450644editdlrm
SparseCsrTensorUtils.h5230644editdlrm
SparseTensorImpl.h124170644editdlrm
SparseTensorUtils.h42190644editdlrm
Storage.h430644editdlrm
Tensor.h480644editdlrm
TensorAccessor.h510644editdlrm
TensorGeometry.h18550644editdlrm
TensorIndexing.h219230644editdlrm
TensorIterator.h299620644editdlrm
TensorIteratorInternal.h18620644editdlrm
TensorMeta.h29170644editdlrm
TensorNames.h25190644editdlrm
TensorOperators.h32750644editdlrm
TensorOptions.h490644editdlrm
TensorUtils.h56870644editdlrm
ThreadLocalState.h32890644editdlrm
TracerMode.h55760644editdlrm
TypeDefault.h6800644editdlrm
Utils.h59930644editdlrm
Version.h3400644editdlrm
VmapMode.h9520644editdlrm
VmapTransforms.h76540644editdlrm
WrapDimUtils.h34380644editdlrm
WrapDimUtilsMulti.h7680644editdlrm
Edit: /usr/local/lib64/python3.6/site-packages/torch/include/ATen/record_function.h (24044B)
#pragma once #include #include #include #include #include #include #include #include #include namespace c10 { class TORCH_API OperatorHandle; } namespace at { // Kind of record function scope; enum class C10_API_ENUM RecordScope : uint8_t { // c10/ATen ops, autograd nodes FUNCTION = 0, // Functions/nodes called from the autograd BACKWARD_FUNCTION, // TorchScript functions, methods TORCHSCRIPT_FUNCTION, // Kernel Function dtype Tag KERNEL_FUNCTION_DTYPE, // Kernel Function dtype Tag LITE_INTERPRETER, // User defined scope (e.g. with record_function()) USER_SCOPE, NUM_SCOPES, // must be the last in the list }; } // namespace at namespace std { template <> struct hash { size_t operator()( const at::RecordScope& sc) const { return static_cast(sc); } }; } // namespace std namespace at { struct TORCH_API StringView { StringView() : StringView(nullptr) {} explicit StringView(const char* str_ptr) : owned_str_ptr_(nullptr), str_ptr_(str_ptr) {} explicit StringView(std::string str) : owned_str_ptr_(std::make_shared(std::move(str))), str_ptr_(owned_str_ptr_->c_str()) {} const char* str() const { return str_ptr_; } friend std::ostream& operator<<(std::ostream& os, const StringView& dt) { os << dt.str(); return os; } friend bool operator==(const StringView& lhs, const StringView& rhs) { return strcmp(lhs.str(), rhs.str()) == 0; } friend bool operator!=(const StringView& lhs, const StringView& rhs) { return !(lhs == rhs); } private: std::shared_ptr owned_str_ptr_; const char* str_ptr_; }; // Soft limit on the number of callbacks to use; constexpr std::size_t kSoftLimitCallbacks = 4; // An abstract base class for various observer contexts that can be attached to // the RecordFunction. struct ObserverContext { virtual ~ObserverContext() {} protected: ObserverContext() {} }; typedef c10::SmallVector CallbackHandles; typedef std::vector> ObserverContextList; typedef uint64_t RecordFunctionHandle; struct TORCH_API RecordFunction { // Default constructor is used with before function called afterwards: // scope - record scope that this function tracks // pre_sampled - whether this RecordFunction was already pre-sampled with // kLowProb probability RecordFunction( RecordScope scope = RecordScope::FUNCTION, bool pre_sampled = false); template void before( F fn, const std::vector* args, int64_t current_sequence_nr = -1) { if (!isActive()) { return; } state_->inputs_ = *args; before(fn, current_sequence_nr); } // Destructor calls end callbacks virtual ~RecordFunction(); RecordFunction(const RecordFunction&) = delete; RecordFunction& operator=(const RecordFunction&) = delete; const StringView& name() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called name() on inactive RecordFunction"); return state_->name_; } int64_t seqNr() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called seqNr() on inactive RecordFunction"); return state_->sequence_nr_; } const std::vector& inputs() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called inputs() on inactive RecordFunction"); return state_->inputs_; } const std::vector& outputs() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called outputs() on inactive RecordFunction"); return state_->outputs_; } void setOutputs(std::vector&& outputs) const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called setOutputs() on inactive RecordFunction"); state_->outputs_ = std::move(outputs); } void setOutputs(c10::ArrayRef outputs) const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called setOutputs() on inactive RecordFunction"); state_->outputs_ = outputs.vec(); } size_t num_inputs() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called num_inputs() on inactive RecordFunction"); return state_->op_input_size; } size_t num_outputs() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called num_outputs() on inactive RecordFunction"); return state_->op_output_size; } // Retrieves the thread_id that this RecordFunction ran start callbacks with. // Useful for writing thread safe end callbacks that may be potentially // executed in a different thread (async ops) uint64_t threadId() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called threadId() on inactive RecordFunction"); return state_->thread_id_; } // For backward functions - thread id of the corresponding forward function, // or zero otherwise; // used alongside with sequence number to correlate backward functions with // the forward ones uint64_t forwardThreadId() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called forwardThreadId() on inactive RecordFunction"); return state_->fwd_thread_id_; } void setForwardThreadId(uint64_t thread_id) { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called setForwardThreadId() on inactive RecordFunction"); state_->fwd_thread_id_ = thread_id; } RecordScope scope() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called scope() on inactive RecordFunction"); return state_->scope_; } // Returns logical thread_id for the current thread static uint64_t currentThreadId(); // Internal functions, do not use directly; // used in python's context manager // before functions initialize RecordFunction members and call // start callbacks void before(const char* name, int64_t sequence_nr = -1); void before(std::string name, int64_t sequence_nr = -1); void before(c10::OperatorHandle const& op, int64_t sequence_nr = -1); // Sets node ID for distributed profiling static void setDefaultNodeId(int64_t defaultNodeId); // Gets node ID for distributed profiling static int64_t getDefaultNodeId(); template void before( F fn, c10::ArrayRef args, int64_t current_sequence_nr = -1) { if (!isActive()) { return; } state_->inputs_ = args.vec(); before(fn, current_sequence_nr); } template void before( F fn, std::vector&& args, int64_t current_sequence_nr = -1) { if (!isActive()) { return; } state_->inputs_ = std::move(args); before(fn, current_sequence_nr); } // Calls end callbacks. After end(), accessors will no longer provide useful results. void end(); // Internal-only, used only force async event for distributed events profiling. void _setAsync(); // Returns whether this RecordFunction corresponds to an async event orn ot. bool isAsync() const; RecordFunctionHandle handle() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called handle() on inactive RecordFunction"); return state_->handle_; } c10::optional operator_name() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called operator_name() on inactive RecordFunction"); return state_->operator_name_; } void setHandle(RecordFunctionHandle handle) { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called setHandle() on inactive RecordFunction"); state_->handle_ = handle; } // Whether this RecordFunction runs any callbacks. bool isActive() const { return state_ != nullptr; } bool needsInputs() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called needsInputs() on inactive RecordFunction"); return state_->needs_inputs; } bool needsOutputs() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called needsOutputs() on inactive RecordFunction"); return state_->needs_outputs; } int64_t debugHandle() const { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called debugHandle() on inactive RecordFunction"); return state_->debug_handle_; } void setDebugHandle(int64_t debug_handle) { TORCH_INTERNAL_ASSERT_DEBUG_ONLY(state_, "Called setDebugHandle() on inactive RecordFunction"); state_->debug_handle_ = debug_handle; } private: // Allows the modification of some internal states for callbacks. friend class CallbackManager; struct State { explicit State(RecordScope scope) : scope_(scope) {} // Whether any of the picked callbacks require inputs bool needs_inputs = false; // Whether any of the picked callbacks require outputs bool needs_outputs = false; // In cases when RecordFunction might be active but we chose not to // use the observers (e.g. operator is not observed), this boolean // flag is used to check whether the start callbacks were called bool called_start_callbacks_ = false; // Whether the RecordFunction is pre-sampled bool pre_sampled_ = false; // Used internally to keep track of thread local and global callbacks // that were picked to run; must be sorted; CallbackHandles sorted_active_tls_handles_; CallbackHandles sorted_active_global_handles_; // Stores various ObserverContext objects with event metadata for thread local // callbacks. ObserverContextList tls_ctx_; // Stores various ObserverContext objects with event metadata for global // callbacks. ObserverContextList global_ctx_; StringView name_; int64_t sequence_nr_ = -1; std::vector inputs_; std::vector outputs_; c10::optional operator_name_; size_t op_input_size{0}; size_t op_output_size{0}; // Kind of scope this RecordFunction is observing const RecordScope scope_; // The logical thread_id that this RecordFunction was created with uint64_t thread_id_ = 0; // For backward functions - thread id of the the forward function uint64_t fwd_thread_id_ = 0; // Unique id for this RecordFunction, used in callbacks to track start // and end of ranges RecordFunctionHandle handle_ {0}; // Whether this record_function corresponds to an async event or not. Async // events can complete in different threads or follow a future-like pattern // of use. bool is_async_{false}; // Debug handles are used for lazy annotation of module hierarchy // and callstack. // This is specifically is useful for mobile runtime, where generated // debug handles can be lazily symbolicated using debug information int64_t debug_handle_{-1}; }; std::unique_ptr state_; }; // // PyTorch callbacks/observers API: // /** * RecordFunctionCallback represents a pair of callbacks to be used with * RecordFunction, members: * start, end - the callbacks to run when entering and exiting the scope; * optionally, the start callback may return an ObserverContext which will * be passed to the end callback, use appropriate constructor accordingly. * needs_inputs - whether the callbacks need the inputs passed from the observed * function/range; NOTE: passing the inputs incurs an additional overhead; * sampling_probability - if not 1.0, then the callback is probabilistically sampled * to run; NOTE: start and end callbacks always run as a pair and are sampled * together; * scopes - types of scopes to execute the callbacks on (see RecordScope); * passing empty set means the callbacks will be executed for all possible * scope types * should_run - optional function that returns whether this callback should run; * overwrites the effect of setting sampling_probability */ class TORCH_API RecordFunctionCallback { public: using StartCallback = std::unique_ptr(*)(const RecordFunction&); using EndCallback = void (*)(const RecordFunction&, ObserverContext*); // This interface supports observers that require passing an ObserverContext // between start and end callbacks. explicit RecordFunctionCallback( StartCallback start, EndCallback end = nullptr) : start_(start), end_(end) { scopes_.fill(true); } RecordFunctionCallback& needsInputs(bool needs_inputs) { needs_inputs_ = needs_inputs; return *this; } RecordFunctionCallback& needsOutputs(bool needs_outputs) { needs_outputs_ = needs_outputs; return *this; } RecordFunctionCallback& needsIds(bool needs_ids) { needs_ids_ = needs_ids; return *this; } RecordFunctionCallback& samplingProb(double sampling_prob) { TORCH_CHECK(sampling_prob >= 0.0 && sampling_prob <= 1.0, "Invalid sampling probability"); sampling_prob_ = sampling_prob; return *this; } RecordFunctionCallback& scopes( const std::unordered_set>& scopes) { if (!scopes.empty()) { scopes_.fill(false); for (auto sc : scopes) { scopes_[static_cast(sc)] = true; } } else { scopes_.fill(true); } return *this; } bool needsInputs() const { return needs_inputs_; } bool needsOutputs() const { return needs_outputs_; } bool needsIds() const { return needs_ids_; } double samplingProb() const { return sampling_prob_; } bool checkScope(RecordScope sc) const { return scopes_[(size_t)sc]; } StartCallback start() const { return start_; } EndCallback end() const { return end_; } private: friend class CallbackManager; StartCallback start_; EndCallback end_; double sampling_prob_ = 1.0; std::array(RecordScope::NUM_SCOPES)> scopes_ = {}; bool needs_inputs_ = false; bool needs_outputs_ = false; bool needs_ids_ = false; }; // Using macro to minimize inputs copies, // optional argument - function's seq_no #define RECORD_FUNCTION_WITH_SCOPE(scope, fn, inputs, ...) \ at::RecordFunction guard(scope); \ if (guard.isActive()) { \ if (guard.needsInputs()) { \ guard.before(fn, inputs, ##__VA_ARGS__); \ } else { \ guard.before(fn, ##__VA_ARGS__); \ } \ } #define RECORD_FUNCTION(fn, inputs, ...) \ RECORD_FUNCTION_WITH_SCOPE( \ at::RecordScope::FUNCTION, \ fn, inputs, ##__VA_ARGS__) #define RECORD_TORCHSCRIPT_FUNCTION(mn, inputs) \ RECORD_FUNCTION_WITH_SCOPE( \ at::RecordScope::TORCHSCRIPT_FUNCTION, mn, inputs) // Custom user scopes in C++; similar to Python's 'with record_function("..."):' #define RECORD_USER_SCOPE(fn) \ RECORD_FUNCTION_WITH_SCOPE( \ at::RecordScope::USER_SCOPE, fn, {}) // RECORD_USER_SCOPE with inputs #define RECORD_USER_SCOPE_WITH_INPUTS(fn, inputs) \ RECORD_FUNCTION_WITH_SCOPE( \ at::RecordScope::USER_SCOPE, fn, inputs) // Helper macro to pass in debug handle that is used to // post process events #define RECORD_WITH_SCOPE_DEBUG_HANDLE_AND_INPUTS( \ scope, fn, debug_handle, inputs, ...) \ at::RecordFunction guard(scope); \ if (guard.isActive()) { \ guard.setDebugHandle(debug_handle); \ if (guard.needsInputs()) { \ guard.before(fn, inputs, ##__VA_ARGS__); \ } else { \ guard.before(fn, ##__VA_ARGS__); \ } \ } // Helper macros to record LITE INTERPETER scope events with debug handles #define RECORD_EDGE_SCOPE_WITH_DEBUG_HANDLE_AND_INPUTS( \ fn, debug_handle, inputs) \ RECORD_WITH_SCOPE_DEBUG_HANDLE_AND_INPUTS( \ at::RecordScope::LITE_INTERPRETER, fn, debug_handle, inputs) // Notes: // - two types of callbacks are provided: thread local and global // - thread local callbacks are added/removed only for the given thread // and are stored locally for each thread and separately from the list // of the global callbacks // - global callbacks are stored in a single per process list and are // invoked by every RecordFunction, in addition to the thread local // callbacks specific to the given thread // - we allow the added callbacks to be sampled, by specifying a sampling // probability for each callback pair, if the start callback is // not picked to run, the corresponding end callback won't be called // - a typical use case for the global callbacks is passive monitoring // in the background (e.g. fleet-wide monitoring), without focusing on // the specific peice of code // - in contrast, thread local callbacks are enabled locally, on demand, // for the specific piece of code (range) and are not sampled // - a typical use case for thread local callbacks is profiler and code // execution tracer // - note, thread local callbacks are automatically propagated with // ThreadLocalState across JIT continuations and async tasks (at::launch) // - adding/removing global callbacks is not thread safe and should be done // only when no other code is running, e.g. during the initialization typedef uint64_t CallbackHandle; struct GlobalRecordFunctionCallbacksEntry { RecordFunctionCallback callback; private: std::atomic enabled; public: CallbackHandle handle; GlobalRecordFunctionCallbacksEntry(RecordFunctionCallback&& cb, CallbackHandle h) : callback(std::move(cb)), enabled(true), handle(h) {} // Copying is fine despite std::atomic not being supposed to // have a copy/move constructor: adding & removing callbacks is // already not thread-safe. GlobalRecordFunctionCallbacksEntry( const GlobalRecordFunctionCallbacksEntry& rhs) : callback(rhs.callback), enabled(rhs.enabled.load()), handle(rhs.handle) {} GlobalRecordFunctionCallbacksEntry& operator=(const GlobalRecordFunctionCallbacksEntry& rhs) { callback = rhs.callback; enabled = rhs.enabled.load(); handle = rhs.handle; return *this; } GlobalRecordFunctionCallbacksEntry( GlobalRecordFunctionCallbacksEntry&& rhs) noexcept : callback(std::move(rhs.callback)), enabled(rhs.enabled.load()), handle(rhs.handle) {} GlobalRecordFunctionCallbacksEntry& operator=(GlobalRecordFunctionCallbacksEntry&& rhs) noexcept { callback = std::move(rhs.callback); enabled = rhs.enabled.load(); handle = rhs.handle; return *this; } // Returns true if the status changed, false otherwise. bool disable() { bool expected = true; // NOTE: we use sequentially consistent access here and in // enable() because updating further atomic flags depends on this // operation. return enabled.compare_exchange_strong(expected, false); } // Returns true if the status changed, false otherwise. bool enable() { bool expected = false; return enabled.compare_exchange_strong(expected, true); } // Read the flag. Note that it is neither necessary nor correct to // check this before calling enable() or disable(). bool isEnabled() const { return enabled.load(std::memory_order_relaxed); } }; // It is unnecessary to use atomic operations for enabling // thread-local function callbacks. Moreover, it prevents saving to // ThreadLocalState because std::atomic is non-copyable. struct ThreadLocalRecordFunctionCallbacksEntry { RecordFunctionCallback callback; bool enabled = true; CallbackHandle handle; ThreadLocalRecordFunctionCallbacksEntry(RecordFunctionCallback&& cb, CallbackHandle h) : callback(std::move(cb)), handle(h) {} bool disable() { auto old = enabled; enabled = false; return old != enabled; } bool enable() { auto old = enabled; enabled = true; return old != enabled; } bool isEnabled() const { return enabled; } }; // Holds pairs (callbacks, unique_id) using GlobalRecordFunctionCallbacks = std::vector; using ThreadLocalRecordFunctionCallbacks = std::vector; /** * addThreadLocalCallback adds a thread local callback to run with RecordFunction, * returns handle to use with removeThreadLocalCallback */ TORCH_API CallbackHandle addThreadLocalCallback( RecordFunctionCallback cb); /** * hasThreadLocalCallbacks returns whether there're callbacks registered * with addThreadLocalCallback */ TORCH_API bool hasThreadLocalCallbacks(); /** * clearThreadLocalCallbacks removes all thread local callbacks */ TORCH_API void clearThreadLocalCallbacks(); /** * addGlobalCallback adds a global callback to run with RecordFunction: * * WARNING: not thread safe, typically addGlobalCallback can be called * only during the program initialization */ TORCH_API CallbackHandle addGlobalCallback( RecordFunctionCallback cb); /** * removeCallback removes a callback given the handle returned by * addThreadLocalCallback or addGlobalCallback; * * WARNING: removing a global callback is not thread safe, * no other code can run simultaneously */ TORCH_API void removeCallback(CallbackHandle handle); /** * Prevent the given callback from executing. If handle is invalid, * does nothing. */ TORCH_API void disableCallback(CallbackHandle handle); /** * Allow the given callback, previously disabled with disableCallback, to * execute again. If handle is invalid, does nothing. */ TORCH_API void reenableCallback(CallbackHandle handle); /** * hasGlobalCallbacks returns whether there're global callbacks * registered with pushGlobalCallback */ TORCH_API bool hasGlobalCallbacks(); /** * clearGlobalCallbacks removes all global callbacks * WARNING: not thread safe */ TORCH_API void clearGlobalCallbacks(); // for both thread local and global callbacks TORCH_API bool hasCallbacks(); TORCH_API void clearCallbacks(); // not thread safe /** * enableRecordFunction enables RecordFunction thread locally */ TORCH_API void enableRecordFunction(bool enable = true); /** * isRecordFunctionEnabled returns whether RecordFunction * is enabled thread locally */ TORCH_API bool isRecordFunctionEnabled(); class TORCH_API RecordFunctionGuard { public: explicit RecordFunctionGuard(bool is_enabled = true) : prev_value_(isRecordFunctionEnabled()) { enableRecordFunction(is_enabled); } virtual ~RecordFunctionGuard() { enableRecordFunction(prev_value_); } private: bool prev_value_ = false; }; class TORCH_API DisableRecordFunctionGuard : public RecordFunctionGuard { public: DisableRecordFunctionGuard() : RecordFunctionGuard(false) {} virtual ~DisableRecordFunctionGuard() {} }; struct TORCH_API RecordFunctionTLS { // Thread local vector of callbacks, holds pairs (callbacks, unique_id); // must be sorted in increasing handles order ThreadLocalRecordFunctionCallbacks sorted_tls_callbacks_; bool tls_record_function_enabled_ = true; // Stores the number of coin flips before the next successful coin flip int tries_left_ = 0; }; TORCH_API const RecordFunctionTLS& get_record_function_tls_(); TORCH_API void set_record_function_tls_(const RecordFunctionTLS& tls); // Checks whether RecordFunction should be called, // sets boolean pointed by the argument to whether pre-sampling was used TORCH_API bool shouldRunRecordFunction(bool*); // The following functions are used to disable/enable pre-sampling of RecordFunction // when high-frequency/non-sampled callbacks are added/removed. // Note: every call to bumpRecordAllFunctions() is supposed to be matched with // the corresponding releaseRecordAllFunctions() call. // Note: disabling pre-sampling of RecordFunction incurs an extra overhead, since // RecordFunction will be created for each operator call. TORCH_API void bumpRecordAllFunctions(); TORCH_API void releaseRecordAllFunctions(); TORCH_API bool checkRecordAllFunctions(); } // namespace at