Testing NetworkMonitor object.
| 46 | |
| 47 | |
| 48 | class TestNetworkMonitor: |
| 49 | """ |
| 50 | Testing NetworkMonitor object. |
| 51 | """ |
| 52 | |
| 53 | network = Network() |
| 54 | |
| 55 | inpt = Input(25) |
| 56 | network.add_layer(inpt, name="X") |
| 57 | _if = IFNodes(75) |
| 58 | network.add_layer(_if, name="Y") |
| 59 | conn = Connection(inpt, _if, w=torch.rand(inpt.n, _if.n)) |
| 60 | network.add_connection(conn, source="X", target="Y") |
| 61 | |
| 62 | mon = NetworkMonitor(network, state_vars=["s", "v", "w"]) |
| 63 | network.add_monitor(mon, name="monitor") |
| 64 | |
| 65 | network.run(inputs={"X": torch.bernoulli(torch.rand(50, inpt.n))}, time=50) |
| 66 | |
| 67 | recording = mon.get() |
| 68 | |
| 69 | assert recording["X"]["s"].size() == torch.Size([50, 1, inpt.n]) |
| 70 | assert recording["Y"]["s"].size() == torch.Size([50, 1, _if.n]) |
| 71 | assert recording["Y"]["s"].size() == torch.Size([50, 1, _if.n]) |
| 72 | |
| 73 | del network.monitors["monitor"] |
| 74 | |
| 75 | mon = NetworkMonitor(network, state_vars=["s", "v", "w"], time=50) |
| 76 | network.add_monitor(mon, name="monitor") |
| 77 | |
| 78 | network.run(inputs={"X": torch.bernoulli(torch.rand(50, inpt.n))}, time=50) |
| 79 | |
| 80 | recording = mon.get() |
| 81 | |
| 82 | assert recording["X"]["s"].size() == torch.Size([50, 1, inpt.n]) |
| 83 | assert recording["Y"]["s"].size() == torch.Size([50, 1, _if.n]) |
| 84 | assert recording["Y"]["s"].size() == torch.Size([50, 1, _if.n]) |
| 85 | |
| 86 | |
| 87 | if __name__ == "__main__": |
no test coverage detected