Runs a single simulation step. :param x: Inputs to the layer.
(self, x: torch.Tensor)
| 498 | ) # Lower bound of voltage. |
| 499 | |
| 500 | def forward(self, x: torch.Tensor) -> None: |
| 501 | # language=rst |
| 502 | """ |
| 503 | Runs a single simulation step. |
| 504 | |
| 505 | :param x: Inputs to the layer. |
| 506 | """ |
| 507 | # Decay voltages. |
| 508 | self.v = self.decay * (self.v - self.rest) + self.rest |
| 509 | |
| 510 | # Integrate inputs. |
| 511 | x.masked_fill_(self.refrac_count > 0, 0.0) |
| 512 | |
| 513 | # Decrement refractory counters. |
| 514 | self.refrac_count -= self.dt |
| 515 | |
| 516 | self.v += x # interlaced |
| 517 | |
| 518 | # Check for spiking neurons. |
| 519 | self.s = self.v >= self.thresh |
| 520 | |
| 521 | # Refractoriness and voltage reset. |
| 522 | self.refrac_count.masked_fill_(self.s, self.refrac) |
| 523 | self.v.masked_fill_(self.s, self.reset) |
| 524 | |
| 525 | # Voltage clipping to lower bound. |
| 526 | if self.lbound is not None: |
| 527 | self.v.masked_fill_(self.v < self.lbound, self.lbound) |
| 528 | |
| 529 | super().forward(x) |
| 530 | |
| 531 | def reset_state_variables(self) -> None: |
| 532 | # language=rst |