File size: 6,305 Bytes
3d46076
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
import numpy as np
from typing import Tuple, List, Dict, Optional
from src.connectome.types import ConnectomeGraph

class PlasticityEngine:
    def __init__(
        self,
        learning_rate: float = 0.05,
        weight_decay: float = 0.01,
        min_weight: float = 0.01,
        max_weight: float = 1.0,
        prune_threshold: float = 0.02,
        rewire_rate: float = 0.01
    ):
        # Defaults match shaders/plasticity.comp + cpu_reference.cpu_plasticity_step
        # so CPU-runtime and GPU-runtime plasticity are the same rule.
        self.learning_rate = learning_rate
        self.weight_decay = weight_decay
        self.min_weight = min_weight
        self.max_weight = max_weight
        self.prune_threshold = prune_threshold
        self.rewire_rate = rewire_rate

    def apply_hebbian_update(
        self,
        graph: ConnectomeGraph,
        pre_activations: Optional[np.ndarray] = None,
        post_activations: Optional[np.ndarray] = None,
        reward: float = 0.0,
        prev_activations: Optional[np.ndarray] = None,
        current_activations: Optional[np.ndarray] = None
    ) -> int:
        """
        Applies reward-modulated Hebbian update directly to CSR weights.
        Returns the number of synapses updated.
        """
        pre = pre_activations if pre_activations is not None else prev_activations
        post = post_activations if post_activations is not None else current_activations
        if pre is None or post is None:
            return 0

        M = len(graph.weights)
        for i in range(graph.num_neurons):
            start = graph.row_offsets[i]
            end = graph.row_offsets[i + 1]
            a_post = float(post[i])
            for k in range(start, end):
                pre_idx = graph.col_indices[k]
                a_pre = float(pre[pre_idx])
                w = graph.weights[k]
                
                # Three-factor rule: pre * post * reward - decay
                delta = self.learning_rate * reward * (a_pre * a_post - self.weight_decay * w)
                graph.weights[k] = float(np.clip(w + delta, self.min_weight, self.max_weight))
        # Provenance: weights changed -> graph hash MUST be refreshed
        # (stale hashes would silently break snapshot/restore equivalence).
        graph.graph_hash = graph.compute_graph_hash()
        return M

    def prune_weak_synapses(self, graph: ConnectomeGraph) -> int:
        """
        Prunes synapses with weight < prune_threshold.
        Rebuilds the CSR arrays in-place.
        Returns number of pruned synapses.
        """
        new_row_offsets = [0]
        new_col_indices = []
        new_weights = []
        pruned_count = 0

        for i in range(graph.num_neurons):
            start = graph.row_offsets[i]
            end = graph.row_offsets[i + 1]
            for k in range(start, end):
                w = graph.weights[k]
                if w >= self.prune_threshold:
                    new_col_indices.append(graph.col_indices[k])
                    new_weights.append(w)
                else:
                    pruned_count += 1
            new_row_offsets.append(len(new_col_indices))

        graph.row_offsets = np.array(new_row_offsets, dtype=np.int32)
        graph.col_indices = np.array(new_col_indices, dtype=np.int32)
        graph.weights = np.array(new_weights, dtype=np.float32)
        graph.graph_hash = graph.compute_graph_hash()
        graph.validate_invariants()
        return pruned_count

    def rewire_coactive_synapses(
        self,
        graph: ConnectomeGraph,
        activations: np.ndarray,
        max_new_synapses: int = 20,
        seed: Optional[int] = None
    ) -> int:
        """
        Creates new synaptic connections between pairs of neurons that are simultaneously active
        but currently disconnected, respecting biological spatial proximity.
        """
        rng = np.random.RandomState(seed)
        active_indices = np.where(activations > 0.4)[0]
        if len(active_indices) < 2:
            return 0

        added_count = 0
        new_edges = {i: [] for i in range(graph.num_neurons)}
        
        # Populate existing edges
        for i in range(graph.num_neurons):
            start = graph.row_offsets[i]
            end = graph.row_offsets[i + 1]
            for k in range(start, end):
                new_edges[i].append((graph.col_indices[k], graph.weights[k]))

        # Search for candidates among active neurons
        # CSR CONVENTION v3: creating directed edge a->b appends source a to row b.
        shuffled = rng.permutation(active_indices)
        for idx_a in shuffled:
            if added_count >= max_new_synapses:
                break
            coord_a = graph.coordinates[idx_a]

            for idx_b in shuffled:
                if idx_a == idx_b:
                    continue
                existing_targets = {s for s, _ in new_edges[idx_b]}
                if idx_a in existing_targets:
                    continue
                coord_b = graph.coordinates[idx_b]
                dist = np.linalg.norm(coord_a - coord_b)
                # Within biological reach (e.g. 8000 nm)
                if dist < 8000.0:
                    init_weight = float(0.1 + 0.1 * rng.rand())
                    new_edges[idx_b].append((int(idx_a), init_weight))
                    added_count += 1
                    if added_count >= max_new_synapses:
                        break

        if added_count > 0:
            # Reconstruct CSR
            new_row_offsets = [0]
            new_col_indices = []
            new_weights = []
            for i in range(graph.num_neurons):
                # Sort targets by index
                new_edges[i].sort(key=lambda x: x[0])
                for target, w in new_edges[i]:
                    new_col_indices.append(target)
                    new_weights.append(w)
                new_row_offsets.append(len(new_col_indices))

            graph.row_offsets = np.array(new_row_offsets, dtype=np.int32)
            graph.col_indices = np.array(new_col_indices, dtype=np.int32)
            graph.weights = np.array(new_weights, dtype=np.float32)
            graph.graph_hash = graph.compute_graph_hash()
            graph.validate_invariants()

        return added_count