FlyBrain-Lab / src /brain /plasticity.py
timfromhcs's picture
FlyBrain v4.1.0 Space build (REAL_SUBGRAPH, CPU-only, honest backend)
3d46076 verified
Raw
History Blame Contribute Delete
6.31 kB
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