@@ -130,6 +130,7 @@ class CacheEntry {
130130 // The caller is expected to check `GlobalCallbackManager::get().version()'
131131 // and call CacheEntry::update() if necessary.
132132 StepCallbacks getActiveCallbacks ();
133+ c10::optional<StepCallbacks> getActiveCallbacksUnlessEmpty ();
133134
134135 // Full rebuild. (E.g. during registration)
135136 void update (const std::vector<RecordFunctionCallback>& callbacks);
@@ -142,6 +143,8 @@ class CacheEntry {
142143 int tries_left_{-1 };
143144 };
144145
146+ C10_ALWAYS_INLINE void getActiveCallbacksImpl ();
147+
145148 void rebuildActiveCallbacks ();
146149 int sampleTries (double p) const ;
147150
@@ -169,6 +172,7 @@ class LocalCallbackManager {
169172 public:
170173 const RecordFunctionTLS& getTLS () const ;
171174 StepCallbacks getActiveCallbacks (const RecordScope scope);
175+ c10::optional<StepCallbacks> getActiveCallbacksUnlessEmpty (const RecordScope scope);
172176
173177 void setTLS (const RecordFunctionTLS& tls);
174178 void seed (uint32_t seed);
@@ -178,6 +182,8 @@ class LocalCallbackManager {
178182 void clearCallbacks ();
179183
180184 private:
185+ void rebuildActiveCallbacksIfNeeded ();
186+
181187 void rebuild_all (const GlobalCallbackManager::snapshot_t & global_snapshot);
182188
183189 void rebuild_callback_scopes (
@@ -271,7 +277,7 @@ void CacheEntry::update(const std::vector<RecordFunctionCallback>& callbacks) {
271277 rebuildActiveCallbacks ();
272278}
273279
274- StepCallbacks CacheEntry::getActiveCallbacks () {
280+ void CacheEntry::getActiveCallbacksImpl () {
275281 // We rebuild the active set when `sampling_countdown_` reaches zero, so if it
276282 // reaches zero at the start of this function something has gone wrong.
277283 TORCH_INTERNAL_ASSERT (sampling_countdown_ > 0 , sampling_countdown_);
@@ -295,7 +301,18 @@ StepCallbacks CacheEntry::getActiveCallbacks() {
295301 }
296302 }
297303 }
304+ }
298305
306+ StepCallbacks CacheEntry::getActiveCallbacks () {
307+ getActiveCallbacksImpl ();
308+ return active_callbacks_;
309+ }
310+
311+ c10::optional<StepCallbacks> CacheEntry::getActiveCallbacksUnlessEmpty () {
312+ getActiveCallbacksImpl ();
313+ if (C10_LIKELY (active_callbacks_.empty ())) {
314+ return c10::nullopt ;
315+ }
299316 return active_callbacks_;
300317}
301318
@@ -365,15 +382,25 @@ const RecordFunctionTLS& LocalCallbackManager::getTLS() const {
365382 return registered_callbacks_;
366383}
367384
368- StepCallbacks LocalCallbackManager::getActiveCallbacks (
369- const RecordScope scope) {
385+ void LocalCallbackManager::rebuildActiveCallbacksIfNeeded () {
370386 const auto global_version = GlobalCallbackManager::get ().version ();
371387 if (C10_UNLIKELY (global_version != global_version_)) {
372388 rebuild_all (GlobalCallbackManager::get ().getSnapshot ());
373389 }
390+ }
391+
392+ StepCallbacks LocalCallbackManager::getActiveCallbacks (
393+ const RecordScope scope) {
394+ rebuildActiveCallbacksIfNeeded ();
374395 return active_callbacks_[static_cast <size_t >(scope)].getActiveCallbacks ();
375396}
376397
398+ c10::optional<StepCallbacks> LocalCallbackManager::getActiveCallbacksUnlessEmpty (
399+ const RecordScope scope) {
400+ rebuildActiveCallbacksIfNeeded ();
401+ return active_callbacks_[static_cast <size_t >(scope)].getActiveCallbacksUnlessEmpty ();
402+ }
403+
377404void LocalCallbackManager::setTLS (const RecordFunctionTLS& tls) {
378405 registered_callbacks_ = tls;
379406 rebuild_all (GlobalCallbackManager::get ().getSnapshot ());
@@ -572,6 +599,10 @@ StepCallbacks getStepCallbacks(RecordScope scope) {
572599 return LocalCallbackManager::get ().getActiveCallbacks (scope);
573600}
574601
602+ c10::optional<StepCallbacks> getStepCallbacksUnlessEmpty (RecordScope scope) {
603+ return LocalCallbackManager::get ().getActiveCallbacksUnlessEmpty (scope);
604+ }
605+
575606const RecordFunctionTLS& get_record_function_tls_ () {
576607 return LocalCallbackManager::get ().getTLS ();
577608}
0 commit comments