Spaces:
Paused
Paused
File size: 6,909 Bytes
5e0b58b | 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 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """Headless-testable 5D-to-2D heatmap projection for Brain 5D."""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
import numpy as np
import numpy.typing as npt
from matplotlib.axes import Axes
from matplotlib.colorbar import Colorbar
from matplotlib.image import AxesImage
from src.core.spatial_index import unpack_coords
if TYPE_CHECKING:
from src.core.network import NeuralNetwork
HeatmapKind = Literal["activity", "weights", "energy"]
@dataclass(frozen=True, slots=True)
class HeatmapData:
"""Numerical heatmap payload independent of Matplotlib rendering."""
values: npt.NDArray[np.float64]
kind: HeatmapKind
label: str
title: str
class HeatmapProjector:
"""Project sparse 5D network state onto the X-Y plane."""
def __init__(self, network: NeuralNetwork, activity_tau_ticks: float = 50.0):
if activity_tau_ticks <= 0.0:
raise ValueError("activity_tau_ticks must be > 0")
self.network = network
self.activity_tau_ticks = float(activity_tau_ticks)
self._shape = (int(network.dimensions[0]), int(network.dimensions[1]))
def build(self, kind: HeatmapKind) -> HeatmapData:
"""Build one finite X-Y projection for the requested metric."""
if kind == "activity":
return HeatmapData(
values=self.activity(),
kind=kind,
label="Recent activity",
title="Activity heatmap (X-Y projection)",
)
if kind == "weights":
return HeatmapData(
values=self.weights(),
kind=kind,
label="Mean incoming weight",
title="Weight heatmap (X-Y projection)",
)
if kind == "energy":
return HeatmapData(
values=self.energy(),
kind=kind,
label="Mean energy",
title="Energy heatmap (X-Y projection)",
)
raise ValueError(f"Unsupported heatmap kind: {kind}")
def activity(self) -> npt.NDArray[np.float64]:
"""Return recent spike activity projected onto X-Y."""
sums, counts = self._empty_accumulators()
current_tick = self.network.current_tick
for neuron_id, neuron in self.network.neurons.items():
x_coord, y_coord, _, _, _ = unpack_coords(neuron_id)
if neuron.last_spike_tick < 0:
value = 0.0
else:
age = max(0, current_tick - neuron.last_spike_tick)
value = float(np.exp(-age / self.activity_tau_ticks))
sums[x_coord, y_coord] += value
counts[x_coord, y_coord] += 1.0
return self._mean(sums, counts)
def weights(self) -> npt.NDArray[np.float64]:
"""Return mean incoming synaptic weight per target, projected onto X-Y."""
per_neuron_sum: dict[int, float] = {}
per_neuron_count: dict[int, int] = {}
for synapses in self.network.synapses.values():
for synapse in synapses:
per_neuron_sum[synapse.target_id] = (
per_neuron_sum.get(synapse.target_id, 0.0) + synapse.weight
)
per_neuron_count[synapse.target_id] = (
per_neuron_count.get(synapse.target_id, 0) + 1
)
sums, counts = self._empty_accumulators()
for neuron_id in self.network.neurons:
x_coord, y_coord, _, _, _ = unpack_coords(neuron_id)
incoming_count = per_neuron_count.get(neuron_id, 0)
value = (
per_neuron_sum[neuron_id] / incoming_count if incoming_count else 0.0
)
sums[x_coord, y_coord] += value
counts[x_coord, y_coord] += 1.0
return self._mean(sums, counts)
def energy(self) -> npt.NDArray[np.float64]:
"""Return mean neuron energy projected onto X-Y."""
sums, counts = self._empty_accumulators()
for neuron_id, neuron in self.network.neurons.items():
x_coord, y_coord, _, _, _ = unpack_coords(neuron_id)
sums[x_coord, y_coord] += neuron.energy
counts[x_coord, y_coord] += 1.0
return self._mean(sums, counts)
def _empty_accumulators(
self,
) -> tuple[npt.NDArray[np.float64], npt.NDArray[np.float64]]:
return np.zeros(self._shape, dtype=float), np.zeros(self._shape, dtype=float)
@staticmethod
def _mean(
sums: npt.NDArray[np.float64],
counts: npt.NDArray[np.float64],
) -> npt.NDArray[np.float64]:
result = np.zeros_like(sums)
np.divide(sums, counts, out=result, where=counts > 0.0)
return result
class HeatmapView:
"""Render projected heatmap data into an existing Matplotlib axis."""
def __init__(self, axis: Axes):
self.axis = axis
self._image: AxesImage | None = None
self._colorbar: Colorbar | None = None
def render(self, data: HeatmapData) -> None:
"""Render or update a heatmap without creating duplicate colorbars."""
display_values = data.values.T
if self._image is None:
self._image = self.axis.imshow( # pyright: ignore[reportUnknownMemberType]
display_values,
origin="lower",
interpolation="nearest",
cmap="hot",
aspect="auto",
)
self._colorbar = self.axis.figure.colorbar(
self._image, ax=self.axis
) # pyright: ignore[reportUnknownMemberType]
else:
self._image.set_data(
display_values
) # pyright: ignore[reportUnknownMemberType]
finite = display_values[np.isfinite(display_values)]
if finite.size:
value_min = float(np.min(finite))
value_max = float(np.max(finite))
if value_min == value_max:
value_max = value_min + 1.0
self._image.set_clim(
value_min, value_max
) # pyright: ignore[reportUnknownMemberType]
self.axis.set_title(data.title) # pyright: ignore[reportUnknownMemberType]
self.axis.set_xlabel("X") # pyright: ignore[reportUnknownMemberType]
self.axis.set_ylabel("Y") # pyright: ignore[reportUnknownMemberType]
if self._colorbar is not None:
self._colorbar.set_label(
data.label
) # pyright: ignore[reportUnknownMemberType]
def clear(self) -> None:
"""Clear rendered state while keeping the caller-owned axis reusable."""
if self._colorbar is not None:
self._colorbar.remove()
self._colorbar = None
if self._image is not None:
self._image.remove()
self._image = None
self.axis.clear()
|