MCPcopy Create free account
hub / github.com/BICLab/SpikingBrain-7B / test_spike_count_binary_lif_node

Function test_spike_count_binary_lif_node

W8ASpike/Int2Spike/test.py:12–54  ·  view source on GitHub ↗

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)
)

Source from the content-addressed store, hash-verified

10)
11
12def 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
56def test_spike_count_ternary_lif_node(
57 x: torch.Tensor = None,

Callers 1

test.pyFile · 0.85

Calls 2

spike_fake_quantFunction · 0.90

Tested by

no test coverage detected