Skip to content

Commit bbafb9a

Browse files
Merge pull request #775 from BindsNET/performance
Performance Improvements
2 parents e6f211e + f3cdad0 commit bbafb9a

10 files changed

Lines changed: 1365 additions & 197 deletions

File tree

‎.gitignore‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,5 +23,6 @@ dist/*
2323
logs/*
2424
.pytest_cache/*
2525
.vscode/*
26-
.claude/
27-
data/*
26+
data/*
27+
.claude/*
28+
.claude

‎bindsnet/learning/MCC_learning.py‎

Lines changed: 189 additions & 119 deletions
Large diffs are not rendered by default.

‎bindsnet/network/topology.py‎

Lines changed: 106 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -443,6 +443,7 @@ def __init__(
443443
pipeline: list = [],
444444
manual_update: bool = False,
445445
traces: bool = False,
446+
sparse_compute: bool = False,
446447
**kwargs,
447448
) -> None:
448449
# language=rst
@@ -456,6 +457,11 @@ def __init__(
456457
:param manual_update: Set to :code:`True` to disable automatic updates (applying learning rules) to connection features.
457458
False by default, updates called after each time step
458459
:param traces: Set to :code:`True` to record history of connection activity (for monitors)
460+
:param sparse_compute: Set to :code:`True` to read only the rows of the effective
461+
weight for source neurons that spiked. A win when few sources are active;
462+
on CUDA it is applied only for large connections
463+
(``source.n * target.n >= 4e6``), where the required device sync pays
464+
for itself. Ignored otherwise.
459465
"""
460466

461467
super().__init__(source, target, device, pipeline, **kwargs)
@@ -464,49 +470,114 @@ def __init__(
464470
if self.traces:
465471
self.activity = None
466472

473+
self.sparse_compute = sparse_compute
474+
475+
# Cached (a_eff, b_sum) for pipelines whose features are all static
476+
# (see AbstractFeature.is_static). Invalidated whenever a feature can
477+
# change: learning updates, normalize, reset, device/dtype moves.
478+
self._fold_cache = None
479+
480+
def _apply(self, fn, recurse=True):
481+
self._fold_cache = None
482+
return super()._apply(fn, recurse)
483+
467484
def compute(self, s: torch.Tensor) -> torch.Tensor:
468485
# language=rst
469486
"""
470-
Compute pre-activations given spikes using connection weights.
471-
472-
:param s: Incoming spikes.
473-
:return: Incoming spikes multiplied by synaptic weights (with or without
474-
decaying spike activation).
475-
"""
487+
Direct incoming spikes through the connection's feature pipeline.
476488
477-
# Change to numeric type (torch doesn't like booleans for matrix ops)
478-
# Note: .float() is an expensive operation. Use as minimally as possible!
479-
# if s.dtype != torch.float32:
480-
# s = s.float()
489+
Each feature's ``compute`` returns its ``[source.n, target.n]`` value; how
490+
it folds is set by the feature's ``op`` (``"mul"`` default, ``"add"``,
491+
``"sub"``). Folding the recurrence (start ``A = 1``, ``B = 0``):
481492
482-
# Prepare broadcast from incoming spikes to all output neurons
483-
# |conn_spikes| = [batch_size, source.n * target.n]
484-
conn_spikes = s.view(s.size(0), self.source.n, 1).repeat(1, 1, self.target.n)
485-
# TODO: ^ This could probably be optimized
493+
* ``mul`` factor ``a``: ``A <- a * A`` and ``B <- a * B``
494+
* ``add`` term ``b``: ``B <- B + b``
495+
* ``sub`` term ``b``: ``B <- B - b``
486496
487-
# Run through pipeline
488-
for f in self.pipeline:
489-
conn_spikes = f.compute(conn_spikes)
497+
``B`` stays ``None`` (unallocated) unless an additive feature is present,
498+
so a purely multiplicative pipeline is exactly the single ``s @ A`` matmul.
490499
491-
# Sum signals for each of the output/terminal neurons
492-
# |out_signal| = [batch_size, target.n]
493-
if conn_spikes.size() != torch.Size([s.size(0), self.source.n, self.target.n]):
494-
if conn_spikes.is_sparse:
495-
conn_spikes = conn_spikes.to_dense()
496-
conn_spikes = conn_spikes.view(s.size(0), self.source.n, self.target.n)
500+
:param s: Incoming spikes, shape ``[batch, *source.shape]``.
501+
:return: Post-synaptic input, shape ``[batch, *target.shape]``.
502+
"""
503+
s = s.view(s.size(0), self.source.n)
497504

498-
if conn_spikes.is_sparse:
499-
out_signal = conn_spikes.to_dense().sum(1)
505+
deferred = None
506+
if self._fold_cache is not None:
507+
a_eff, b_sum = self._fold_cache
508+
else:
509+
# running product of multiplicative factors, [source.n, target.n]
510+
a_eff = None
511+
# running additive offset, [source.n, target.n]; None while still zero
512+
b_eff = None
513+
for f in self.pipeline:
514+
factor = f.compute(s)
515+
# Compute-time side effects (e.g. per-time-step weight
516+
# normalization) run after the fold has consumed this value.
517+
d = getattr(f, "defer", None)
518+
if d is not None:
519+
deferred = [d] if deferred is None else deferred + [d]
520+
if factor is None:
521+
# Side-effect-only pipeline entries (sub-features) fold as identity.
522+
continue
523+
if isinstance(factor, torch.Tensor) and factor.is_sparse:
524+
# Sparse feature values carry a leading batch dim ([1, src, tgt])
525+
# from prime_feature; densify to the fold's [src, tgt] shape.
526+
factor = factor.to_dense().view(self.source.n, self.target.n)
527+
op = getattr(f, "op", "mul")
528+
if op == "mul":
529+
a_eff = factor if a_eff is None else a_eff * factor
530+
if b_eff is not None:
531+
b_eff = b_eff * factor
532+
else: # additive contribution: "add" -> +factor, "sub" -> -factor
533+
term = factor if op == "add" else -factor
534+
b_eff = term if b_eff is None else b_eff + term
535+
536+
if a_eff is None:
537+
# Degenerate pipeline with no multiplicative feature: every source
538+
# neuron contributes with unit weight.
539+
a_eff = torch.ones(self.source.n, self.target.n, device=s.device)
540+
if not torch.is_floating_point(a_eff):
541+
a_eff = a_eff.float()
542+
543+
# Additive terms apply to every synapse regardless of spikes, so
544+
# their contribution is the source-sum, a constant [target.n] row.
545+
b_sum = b_eff.sum(dim=0) if b_eff is not None else None
546+
547+
if all(getattr(f, "is_static", True) for f in self.pipeline):
548+
self._fold_cache = (a_eff, b_sum)
549+
550+
# The gather pays for its ``nonzero()`` only where that sync is expensive
551+
# relative to the matmul: on CUDA it needs a large weight matrix to win;
552+
# on CPU there is no device sync, so it helps even at small sizes.
553+
use_gather = self.sparse_compute and (
554+
not a_eff.is_cuda or self.source.n * self.target.n >= 4_000_000
555+
)
556+
if use_gather:
557+
# Read only the rows of A for source neurons that spiked this step
558+
# (numerically identical; faster only when few are active).
559+
active = s.any(dim=0).nonzero(as_tuple=False).squeeze(1)
560+
if active.numel() == 0:
561+
out = torch.zeros(
562+
s.size(0), a_eff.size(-1), device=a_eff.device, dtype=a_eff.dtype
563+
)
564+
else:
565+
out = s[:, active].to(a_eff.dtype) @ a_eff.index_select(0, active)
500566
else:
501-
out_signal = conn_spikes.sum(1)
567+
out = s.to(a_eff.dtype) @ a_eff
502568

503-
if self.traces:
504-
self.activity = out_signal
569+
if b_sum is not None:
570+
out = out + b_sum
505571

506-
if out_signal.size() != torch.Size([s.size(0)] + self.target.shape):
507-
return out_signal.view(s.size(0), *self.target.shape)
508-
else:
509-
return out_signal
572+
if deferred is not None:
573+
for fn in deferred:
574+
fn()
575+
576+
if self.traces:
577+
self.activity = out
578+
if out.size() != torch.Size([s.size(0)] + self.target.shape):
579+
return out.view(s.size(0), *self.target.shape)
580+
return out
510581

511582
def compute_window(self, s: torch.Tensor) -> torch.Tensor:
512583
# language=rst
@@ -544,6 +615,7 @@ def update(self, **kwargs) -> None:
544615
learning = kwargs.get("learning", False)
545616
if learning and not self.manual_update:
546617
# Pipeline learning
618+
self._fold_cache = None
547619
for f in self.pipeline:
548620
f.update(**kwargs)
549621

@@ -553,6 +625,7 @@ def normalize(self) -> None:
553625
Normalize all features in the connection.
554626
"""
555627
# Normalize pipeline features
628+
self._fold_cache = None
556629
for f in self.pipeline:
557630
f.normalize()
558631

@@ -563,6 +636,7 @@ def reset_state_variables(self) -> None:
563636
"""
564637
super().reset_state_variables()
565638

639+
self._fold_cache = None
566640
for f in self.pipeline:
567641
f.reset_state_variables()
568642

0 commit comments

Comments
 (0)