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