Runs a single simulation step. :param x: Inputs to the layer.
(self, x: torch.Tensor)
| 760 | self.lbound = lbound # Lower bound of voltage. |
| 761 | |
| 762 | def forward(self, x: torch.Tensor) -> None: |
| 763 | # language=rst |
| 764 | """ |
| 765 | Runs a single simulation step. |
| 766 | |
| 767 | :param x: Inputs to the layer. |
| 768 | """ |
| 769 | # Decay voltages and current. |
| 770 | self.v = self.decay * (self.v - self.rest) + self.rest |
| 771 | self.i *= self.i_decay |
| 772 | |
| 773 | # Decrement refractory counters. |
| 774 | self.refrac_count -= self.dt |
| 775 | |
| 776 | # Integrate inputs. |
| 777 | self.i += x |
| 778 | self.v += (self.refrac_count <= 0).float() * self.i |
| 779 | |
| 780 | # Check for spiking neurons. |
| 781 | self.s = self.v >= self.thresh |
| 782 | |
| 783 | # Refractoriness and voltage reset. |
| 784 | self.refrac_count.masked_fill_(self.s, self.refrac) |
| 785 | self.v.masked_fill_(self.s, self.reset) |
| 786 | |
| 787 | # Voltage clipping to lower bound. |
| 788 | if self.lbound is not None: |
| 789 | self.v.masked_fill_(self.v < self.lbound, self.lbound) |
| 790 | |
| 791 | super().forward(x) |
| 792 | |
| 793 | def reset_state_variables(self) -> None: |
| 794 | # language=rst |