Tests all stable groups of neurons / nodes.
| 13 | |
| 14 | |
| 15 | class TestNodes: |
| 16 | """ |
| 17 | Tests all stable groups of neurons / nodes. |
| 18 | """ |
| 19 | |
| 20 | def test_init(self): |
| 21 | network = Network() |
| 22 | for i, nodes in enumerate( |
| 23 | [Input, McCullochPitts, IFNodes, LIFNodes, AdaptiveLIFNodes, SRM0Nodes] |
| 24 | ): |
| 25 | for n in [1, 100, 10000]: |
| 26 | layer = nodes(n) |
| 27 | network.add_layer(layer=layer, name=f"{i}_{n}") |
| 28 | |
| 29 | assert layer.n == n |
| 30 | assert (layer.s.float() == torch.zeros(n)).all() |
| 31 | |
| 32 | if nodes in [LIFNodes, AdaptiveLIFNodes]: |
| 33 | assert (layer.v == layer.rest * torch.ones(n)).all() |
| 34 | |
| 35 | layer = nodes(n, traces=True, tc_trace=1e5) |
| 36 | network.add_layer(layer=layer, name=f"{i}_traces_{n}") |
| 37 | |
| 38 | assert layer.n == n |
| 39 | assert layer.tc_trace == 1e5 |
| 40 | assert (layer.s.float() == torch.zeros(n)).all() |
| 41 | assert (layer.x == torch.zeros(n)).all() |
| 42 | assert (layer.x == torch.zeros(n)).all() |
| 43 | |
| 44 | if nodes in [LIFNodes, AdaptiveLIFNodes, SRM0Nodes]: |
| 45 | assert (layer.v == layer.rest * torch.ones(n)).all() |
| 46 | |
| 47 | for nodes in [LIFNodes, AdaptiveLIFNodes]: |
| 48 | for n in [1, 100, 10000]: |
| 49 | layer = nodes( |
| 50 | n, rest=0.0, reset=-10.0, thresh=10.0, refrac=3, tc_decay=1.5e3 |
| 51 | ) |
| 52 | network.add_layer(layer=layer, name=f"{i}_params_{n}") |
| 53 | |
| 54 | assert layer.rest == 0.0 |
| 55 | assert layer.reset == -10.0 |
| 56 | assert layer.thresh == 10.0 |
| 57 | assert layer.refrac == 3 |
| 58 | assert layer.tc_decay == 1.5e3 |
| 59 | assert (layer.s.float() == torch.zeros(n)).all() |
| 60 | assert (layer.v == layer.rest * torch.ones(n)).all() |
| 61 | |
| 62 | def test_transfer(self): |
| 63 | if not torch.cuda.is_available(): |
| 64 | return |
| 65 | |
| 66 | for nodes in Nodes.__subclasses__(): |
| 67 | layer = nodes(10) |
| 68 | |
| 69 | layer.to(torch.device("cuda:0")) |
| 70 | |
| 71 | layer_tensors = [ |
| 72 | k for k, v in layer.state_dict().items() if isinstance(v, torch.Tensor) |