Skip to content

Switch spike tensors from uint8 to bool dtype - #801

Open
md-Yusha wants to merge 1 commit into
BindsNET:masterfrom
md-Yusha:switch-spikes-to-bool-tensors
Open

md-Yusha wants to merge 1 commit into
BindsNET:masterfrom
md-Yusha:switch-spikes-to-bool-tensors

Conversation

@md-Yusha

Copy link
Copy Markdown

torch has supported a native bool dtype since 1.2 (2019), but Nodes.s and its downstream consumers were still allocated as uint8. This meant spike checks relied on comparison hacks like (s > 0) instead of direct boolean semantics, and precluded using logical ops (&, |, ~, .any(), .all()) directly on spike tensors.

Changes:

  • Nodes.init / reset_state_variables: self.s now allocated with dtype=torch.bool
  • Learning rules (learning.py): explicit .float() casts added where spikes are used in arithmetic weight updates
  • Encoders (encoding.py): output dtype aligned to bool at the spike-train boundary

Closes #318

torch has supported a native bool dtype since 1.2 (2019), but
Nodes.s and its downstream consumers were still allocated as
uint8. This meant spike checks relied on comparison hacks like
(s > 0) instead of direct boolean semantics, and precluded using
logical ops (&, |, ~, .any(), .all()) directly on spike tensors.

Changes:
- Nodes.__init__ / reset_state_variables: self.s now allocated
  with dtype=torch.bool
- Learning rules (learning.py): explicit .float() casts added
  where spikes are used in arithmetic weight updates
- Encoders (encoding.py): output dtype aligned to bool at the
  spike-train boundary

Closes BindsNET#318
@Hananel-Hazan

Copy link
Copy Markdown
Collaborator

Hi @md-Yusha,

Thank you for picking this up and for helping with BindsNET!

There are a few changes I would like before we merge:

  1. Input.forward (nodes.py line 235): please drop the .bool() and keep self.s = x. Input has to pass real values through. In 2019 RealInput was merged into Input (commit 4ea63ad), and both ann_to_snn and RepeatEncoder feed real values into it. With the cast every non-zero value becomes 1. I converted the network from test/conversion/test_conversion.py after one epoch on MNIST (96.2% on 500 test images): the SNN gets 91.4% on master and 78.2% with the cast. Without the cast your branch passes all tests and gives bit-identical results to master.

  2. Tests: every .byte() input in the tests became .bool(), so nothing tests uint8 input anymore, and users will keep passing it. Please add a test to test/network/test_perf_equivalence.py that runs the same seeded network with uint8, float and bool input and compares the weights with torch.equal. Please also add one that checks real values pass through Input unchanged, since nothing tests that path today.

  3. EnvironmentPipeline.step_ (environment_pipeline.py line 229): when the environment has an encoder and overlay_input > 1, self.overlay_last_obs - obs now fails, because bool tensors cannot be subtracted. On master it was already wrong, uint8 0 - 1 wraps around to 255. Please add obs = obs.float() at the start of the if self.overlay_t > 1: block. I checked it with a Breakout environment, the overlay then holds only 0 and 1, and all tests pass.

  4. CSRMNodes.last_spikes (nodes.py line 1433): please revert this hunk and keep the comment. This buffer is not bool at runtime, set_batch_size replaces it with a float32 tensor (line 1518) because the einsum in forward needs float. A BoolTensor placeholder points to the wrong type.

  5. topology.py line 1566 still says :param ByteTensor mask:, the same fix you did at line 120.

  6. The encoders now return torch.bool instead of torch.uint8, so this can break user code: 1 - spikes raises an error now, and a + b gives a logical OR instead of a count. Please record it in CHANGELOG.md. The top section of that file is ## [Unreleased] — that is where changes go before they are part of a numbered release. Add this line under its ### Changed heading:

    • Encoders now return torch.bool tensors instead of torch.uint8. Arithmetic on spike tensors changes: 1 - spikes raises an error, and a + b is a logical OR rather than a sum.

    CONTRIBUTING.md line 29 asks for this on every pull request.

Don't worry about the black and isort failures, we will run them on our side.

Let me know what you think.

@md-Yusha

Copy link
Copy Markdown
Author

Thanks for the detailed review, @Hananel-Hazan — the conversion test numbers on Input.forward are a great catch, I hadn't thought about RealInput being merged in back in 2019.

Here's what I'll fix:

  • Input.forward: dropping the .bool() cast, keeping self.s = x so real values pass through untouched.
  • Tests: adding a torch.equal check in test_perf_equivalence.py across uint8/float/bool input, plus a test confirming real values survive Input unchanged.
  • EnvironmentPipeline.step_: adding obs = obs.float() at the top of the overlay_t > 1 block — also TIL uint8 was silently wrapping around on master there.
  • CSRMNodes.last_spikes: reverting that hunk, you're right that set_batch_size overwrites it as float32 for the einsum.
  • topology.py:1566: updating the ByteTensor docstring to match.
  • CHANGELOG.md: adding your note under ## [Unreleased] → ### Changed.

Leaving black/isort as-is per your note. Will push an updated commit soon.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Switch spikes from ByteTensors to BoolTensors

2 participants