Test function for SpikeCountBinaryLIFNode. Args: x (torch.Tensor, optional): Input tensor with non-negative integer values. If None, a random tensor will be generated. high (int): Upper bound for random generation (exclusive). siz
(
x: torch.Tensor = None,
low: int = 0,
high: int = 15,
size: tuple = (12, 1024, 2048)
)
| 10 | ) |
| 11 | |
| 12 | def test_spike_count_binary_lif_node( |
| 13 | x: torch.Tensor = None, |
| 14 | low: int = 0, |
| 15 | high: int = 15, |
| 16 | size: tuple = (12, 1024, 2048) |
| 17 | ) -> bool: |
| 18 | """ |
| 19 | Test function for SpikeCountBinaryLIFNode. |
| 20 | |
| 21 | Args: |
| 22 | x (torch.Tensor, optional): Input tensor with non-negative integer values. |
| 23 | If None, a random tensor will be generated. |
| 24 | high (int): Upper bound for random generation (exclusive). |
| 25 | size (tuple): Shape of randomly generated tensor. |
| 26 | |
| 27 | Returns: |
| 28 | bool: True if the spike sequence sums match the original input; False otherwise. |
| 29 | """ |
| 30 | if x is None: |
| 31 | x = torch.randint(low=low, high=high+1, size=size) |
| 32 | else: |
| 33 | if not isinstance(x, torch.Tensor): |
| 34 | raise TypeError("Input x must be a torch.Tensor") |
| 35 | |
| 36 | if not torch.all(x >= 0): |
| 37 | raise ValueError("Input x must be non-negative") |
| 38 | |
| 39 | if not torch.allclose(x, x.round()): |
| 40 | raise ValueError("Input x must contain integer values only") |
| 41 | |
| 42 | x_zero = torch.zeros_like(x, dtype=torch.float32) |
| 43 | |
| 44 | lif = SpikeCountBinaryLIFNode() |
| 45 | spike_sum = spike_fake_quant(x, lif, x_zero) |
| 46 | |
| 47 | match = torch.allclose((x + x_zero).to(dtype=spike_sum.dtype), spike_sum, rtol=1e-3, atol=1e-3) |
| 48 | |
| 49 | if match: |
| 50 | print(f"\u2714 Numerical match: {lif.__class__.__name__} spike sum matches input spike counts.") |
| 51 | else: |
| 52 | print(f"\u2718 Mismatch: {lif.__class__.__name__} spike sum does not match input spike counts.") |
| 53 | |
| 54 | return match |
| 55 | |
| 56 | def test_spike_count_ternary_lif_node( |
| 57 | x: torch.Tensor = None, |
no test coverage detected