From 69528d36793a2eb21751b3a42bde4fc4226f8888 Mon Sep 17 00:00:00 2001 From: Saachi Goyal <156711741+saachigoyall@users.noreply.github.com> Date: Wed, 29 Jul 2026 21:03:22 -0400 Subject: [PATCH 1/2] Fix MSTDPET.reset_state_variables to clear p_plus, p_minus, and moving-average buffer --- bindsnet/learning/MCC_learning.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/bindsnet/learning/MCC_learning.py b/bindsnet/learning/MCC_learning.py index 33ebacd0a..96d8e69b0 100644 --- a/bindsnet/learning/MCC_learning.py +++ b/bindsnet/learning/MCC_learning.py @@ -731,8 +731,13 @@ def _connection_update(self, **kwargs) -> None: ) super().update() - + def reset_state_variables(self) -> None: self.eligibility.zero_() self.eligibility_trace.zero_() + self.p_plus.zero_() + self.p_minus.zero_() + if self.average_update > 0: + self.average_buffer.zero_() + self.average_buffer_index = 0 return From 3af9aa8b4801d94c0afc7c2babe9c7ce97802f21 Mon Sep 17 00:00:00 2001 From: Saachi Goyal <156711741+saachigoyall@users.noreply.github.com> Date: Wed, 29 Jul 2026 21:07:37 -0400 Subject: [PATCH 2/2] Add test for MSTDPET reset_state_variables fix --- test/network/test_learning.py | 49 +++++++++++++++++++++++++++++++++++ 1 file changed, 49 insertions(+) diff --git a/test/network/test_learning.py b/test/network/test_learning.py index 926b921a4..130917a59 100644 --- a/test/network/test_learning.py +++ b/test/network/test_learning.py @@ -268,3 +268,52 @@ def test_rmax(self): time=250, reward=1.0, ) + + def test_mstdpet_reset_clears_moving_average_buffer(self): + # MCC_MSTDPET (MulticompartmentConnection) test + from bindsnet.learning.MCC_learning import MSTDPET as MCC_MSTDPET + from bindsnet.network.topology import MulticompartmentConnection + from bindsnet.network.topology_features import Weight + + network = Network(dt=1.0) + network.add_layer(Input(n=10, traces=True), name="input") + network.add_layer(LIFNodes(n=10, traces=True), name="output") + + weight = Weight( + "weight", + torch.rand(10, 10), + range=[0.0, 1.0], + nu=(1e-2, 1e-2), + learning_rule=MCC_MSTDPET, + tc_plus=20.0, + tc_minus=20.0, + average_update=5, + continues_update=True, + ) + connection = MulticompartmentConnection( + source=network.layers["input"], + target=network.layers["output"], + pipeline=[weight], + ) + network.add_connection(connection, source="input", target="output") + + network.run( + inputs={"input": torch.bernoulli(torch.rand(250, 10)).byte()}, + time=250, + reward=1.0, + ) + + rule = connection.pipeline[0].update_rule + + # sanity check: after 250 steps of activity + reward, state should + # be non-zero before reset + assert rule.average_buffer.abs().sum() > 0 or rule.average_buffer_index != 0 + + rule.reset_state_variables() + + assert torch.all(rule.eligibility == 0) + assert torch.all(rule.eligibility_trace == 0) + assert torch.all(rule.p_plus == 0) + assert torch.all(rule.p_minus == 0) + assert torch.all(rule.average_buffer == 0) + assert rule.average_buffer_index == 0