From 3d4acd2b66d2662a6f1af1c8c8c9b4218dfb880a Mon Sep 17 00:00:00 2001 From: Kael Dai Date: Mon, 13 Jul 2026 18:57:25 -0700 Subject: [PATCH] adding reset option for workshop --- bmtk/simulator/dpointnet/__init__.py | 5 +++++ bmtk/simulator/dpointnet/id_maps.py | 6 ++++++ 2 files changed, 11 insertions(+) diff --git a/bmtk/simulator/dpointnet/__init__.py b/bmtk/simulator/dpointnet/__init__.py index c6fb87fe..1a4fb03d 100644 --- a/bmtk/simulator/dpointnet/__init__.py +++ b/bmtk/simulator/dpointnet/__init__.py @@ -8,3 +8,8 @@ from .input_modules import InputModules from .loss_functions import register_loss_module, add_loss_module from .tf_utils import cleanup_tensorflow, enable_gpu_memory_growth +from .id_maps import TFIDMap + + +def reset(): + TFIDMap().reset() \ No newline at end of file diff --git a/bmtk/simulator/dpointnet/id_maps.py b/bmtk/simulator/dpointnet/id_maps.py index f1aa2a45..53acf1e7 100644 --- a/bmtk/simulator/dpointnet/id_maps.py +++ b/bmtk/simulator/dpointnet/id_maps.py @@ -101,3 +101,9 @@ def tf2bmtk_id_map(self, populations=None): ret_df = pd.concat([ret_df, tmp_df]) return ret_df.set_index('tf_ids') + + def reset(self): + self._bmtk_populations = {} + self._recurrent_tf_indices = [0] + self._recurrent_populations = [] + self._initialized = True