File size: 15,669 Bytes
31226fd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
"""Deterministic end-to-end learning experiment for MHRN.

The experiment demonstrates a complete causal chain:

PRE spikes -> POST spike -> eligibility -> reward -> weight update -> changed response.

It intentionally lives outside the reference core and uses only public network and
learning APIs. The trained weights are evaluated in a fresh network so the reported
response change cannot be explained by residual neuron state.
"""

from __future__ import annotations

import argparse
import itertools
import random
import statistics
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any, cast

import yaml

from src.core import NeuralNetwork
from src.learning.learning_engine import LearningEngine

Config = Mapping[str, Any]
Coord5D = tuple[int, int, int, int, int]
TrialPartitions = dict[str, tuple[int, ...]]


@dataclass(frozen=True, slots=True)
class LearningExperimentResult:
    """Summary of one deterministic system-level learning experiment."""

    training_trials: int
    presynaptic_neurons: int
    initial_mean_weight: float
    final_mean_weight: float
    mean_weight_delta: float
    rewards_received: int
    rewards_applied: int
    reward_weight_updates: int
    baseline_target_spiked: bool
    trained_target_spiked: bool
    baseline_target_peak_v: float
    trained_target_peak_v: float
    baseline_target_spike_tick: int | None
    trained_target_spike_tick: int | None
    train_trial_count: int
    validation_trial_count: int
    holdout_trial_count: int
    protocol_id: str
    protocol_version: int
    condition: str = "learning_on"
    # Partition counts declare a design, not executed validation episodes.
    partition_counts_are_declared: bool = True
    validation_episodes_executed: int = 0
    holdout_episodes_executed: int = 0
    baseline_probes_executed: int = 1
    post_training_probes_executed: int = 1

    @property
    def learned(self) -> bool:
        """Return whether training strengthened weights and changed target response."""
        return (
            self.final_mean_weight > self.initial_mean_weight
            and not self.baseline_target_spiked
            and self.trained_target_spiked
        )


def _experiment_config(config: Config) -> dict[str, Any]:
    """Extract the learning_experiment section from the configuration."""
    section = config.get("learning_experiment", {})
    if not isinstance(section, Mapping):
        raise TypeError("learning_experiment config must be a mapping")
    # Cast to dict[str, Any] to satisfy the type checker
    return cast(dict[str, Any], section)


def _validated_dimensions(config: Config) -> Coord5D:
    """Validate and extract dimensions from the configuration."""
    raw = config.get("dimensions")
    if not isinstance(raw, Sequence):
        raise ValueError("dimensions must be a sequence")

    # Cast to Sequence[int] for type safety
    dims_seq = cast(Sequence[int], raw)
    if len(dims_seq) != 5:
        raise ValueError("dimensions must contain exactly five entries")

    dims = tuple(int(v) for v in dims_seq)
    if len(dims) != 5 or any(v <= 0 for v in dims):
        raise ValueError("all dimensions must be > 0")
    return dims  # pyright: ignore[return-value]


def _candidate_coords(dims: Coord5D) -> Iterable[Coord5D]:
    """Generate all possible 5D coordinates within the given dimensions."""
    product = itertools.product(*(range(size) for size in dims))
    return (cast(Coord5D, coord) for coord in product)


def _validated_trial_partitions(config: Config) -> TrialPartitions:
    """Validate the canonical train/validation/holdout trial split."""
    exp = _experiment_config(config)
    protocol_id = exp.get("protocol_id")
    protocol_version = exp.get("protocol_version")
    if not isinstance(protocol_id, str) or not protocol_id.strip():
        raise ValueError("learning_experiment.protocol_id must not be empty")
    if not isinstance(protocol_version, int) or isinstance(protocol_version, bool):
        raise ValueError("learning_experiment.protocol_version must be an integer")
    if protocol_version < 1:
        raise ValueError("learning_experiment.protocol_version must be positive")
    trials = int(exp.get("training_trials", 20))
    raw = exp.get("partitions")
    if not isinstance(raw, Mapping):
        raise ValueError("learning_experiment.partitions must be a mapping")

    partitions: TrialPartitions = {}
    expected = set(range(trials))
    seen: set[int] = set()
    for name in ("train", "validation", "holdout"):
        values = raw.get(name)
        if not isinstance(values, Sequence) or isinstance(values, (str, bytes)):
            raise ValueError(f"learning_experiment.partitions.{name} must be a list")
        indices = tuple(int(value) for value in values)
        if not indices:
            raise ValueError(f"learning_experiment.partitions.{name} must not be empty")
        if any(index < 0 or index >= trials for index in indices):
            raise ValueError(
                f"learning_experiment.partitions.{name} has out-of-range trial"
            )
        if len(set(indices)) != len(indices) or seen.intersection(indices):
            raise ValueError("learning_experiment partitions must be disjoint")
        seen.update(indices)
        partitions[name] = indices

    if seen != expected:
        raise ValueError(
            "learning_experiment partitions must cover every training trial exactly once"
        )
    return partitions


def _build_convergent_network(
    config: Config,
    weight: float,
) -> tuple[NeuralNetwork, tuple[int, ...], int]:
    """Build a convergent network with presynaptic neurons connected to a target."""
    exp = _experiment_config(config)
    pre_count = int(exp.get("presynaptic_neurons", 48))
    if pre_count <= 0:
        raise ValueError("learning_experiment.presynaptic_neurons must be > 0")

    dims = _validated_dimensions(config)
    target_coord = cast(Coord5D, tuple(size - 1 for size in dims))
    available = [coord for coord in _candidate_coords(dims) if coord != target_coord]
    if pre_count > len(available):
        raise ValueError("not enough coordinates for requested presynaptic neurons")

    # Convert to plain dict for NeuralNetwork constructor
    network_config = dict(config)

    network = NeuralNetwork(network_config, random.Random(int(config.get("seed", 42))))
    pre_ids = tuple(network.add_neuron(coord) for coord in available[:pre_count])
    target_id = network.add_neuron(target_coord)
    delay = int(exp.get("connection_delay_ticks", 1))
    for pre_id in pre_ids:
        network.connect(pre_id, target_id, float(weight), delay)
    network.output_cells.add(target_id)
    return network, pre_ids, target_id


def _advance_to_tick(network: NeuralNetwork, tick: int) -> None:
    """Advance the network to a specific tick."""
    if tick < network.current_tick:
        raise ValueError("cannot move network backwards in time")
    while network.current_tick < tick:
        network.step()


def _reset_trial_dynamics(network: NeuralNetwork) -> None:
    """Reset transient neuron/event state while preserving learned weights.

    Learning trials are declared independent timing episodes. Previously only
    the learning traces were reset, leaving refractory/adaptation state from
    the preceding task and causing valid lower-drive trials to fail.
    """
    network.current_tick = 0
    network.total_spikes = 0
    network.total_events_processed = 0
    network.pending_currents.clear()
    network.event_slots = [[] for _ in range(network.max_delay + 1)]
    network._queued_event_count = 0
    for neuron in network.neurons.values():
        neuron.v = neuron.c
        neuron.u = neuron.b * neuron.v
        neuron.spike_counter = 0
        neuron.last_spike_tick = -1
        neuron.threshold_adaptation = 0.0
        neuron.last_external_current = 0.0
        neuron.last_synaptic_current = 0.0
        neuron.pre_trace = 0.0
        neuron.post_trace = 0.0
        neuron.firing_rate_estimate = 0.0
        neuron._spike_count_window = 0
        neuron._last_update_tick = 0


def _train(
    config: Config, condition: str
) -> tuple[tuple[float, ...], LearningEngine, TrialPartitions]:
    """Train the network using reward-modulated STDP."""
    if condition not in {"learning_on", "learning_off", "sham_replay"}:
        raise ValueError(f"Unsupported learning condition: {condition}")
    exp = _experiment_config(config)
    partitions = _validated_trial_partitions(config)
    reset_trial_dynamics = bool(exp.get("reset_trial_dynamics", False))
    trials = int(exp.get("training_trials", 20))
    spacing = int(exp.get("trial_spacing_ticks", 25))
    pair_delay = int(exp.get("pair_delay_ticks", 5))
    drive = float(exp.get("drive_current", 100.0))
    reward_value = float(exp.get("reward_value", 1.0))
    initial_weight = float(exp.get("initial_weight", 0.05))

    if trials <= 0:
        raise ValueError("learning_experiment.training_trials must be > 0")
    if pair_delay <= 0:
        raise ValueError("learning_experiment.pair_delay_ticks must be > 0")
    if spacing <= pair_delay:
        raise ValueError("trial_spacing_ticks must be greater than pair_delay_ticks")

    training_config = dict(config)
    if condition == "learning_off":
        training_config["eligibility"] = {
            **dict(cast(Mapping[str, Any], config.get("eligibility", {}))),
            "enabled": False,
        }
        training_config["reward"] = {
            **dict(cast(Mapping[str, Any], config.get("reward", {}))),
            "enabled": False,
        }

    network, pre_ids, target_id = _build_convergent_network(
        training_config, initial_weight
    )
    learning = LearningEngine(network, training_config)
    if condition != "learning_off" and not learning.params.reward_enabled:
        raise ValueError("learning experiment requires reward.enabled=true")
    learning.attach()

    for trial in partitions["train"]:
        if reset_trial_dynamics:
            _reset_trial_dynamics(network)
            learning.reset_state()
        pre_tick = trial * spacing
        post_tick = pre_tick + pair_delay
        _advance_to_tick(network, pre_tick)
        for pre_id in pre_ids:
            network.inject_current(pre_id, drive)
        pre_result = network.step()
        if not set(pre_ids).issubset(pre_result.spike_ids):
            raise RuntimeError("training drive failed to spike all presynaptic neurons")

        _advance_to_tick(network, post_tick)
        network.inject_current(target_id, drive)
        post_result = network.step()
        if target_id not in post_result.spike_ids:
            raise RuntimeError("training drive failed to spike target neuron")

        if condition == "sham_replay":
            learning.reset_state()
        learning.set_reward(reward_value, post_result.tick)
        # Each trial is an independent timing episode. Weight changes persist,
        # while timing/eligibility state is cleared to avoid cross-trial pairing.
        learning.reset_state()

    weights = tuple(
        synapse.weight for pre_id in pre_ids for synapse in network.synapses[pre_id]
    )
    return weights, learning, partitions


def _probe_response(
    config: Config,
    weights: Sequence[float],
) -> tuple[bool, float, int | None]:
    """Probe the network response with given weights."""
    exp = _experiment_config(config)
    drive = float(exp.get("drive_current", 100.0))
    probe_ticks = int(exp.get("probe_ticks", 5))
    if probe_ticks < 2:
        raise ValueError("learning_experiment.probe_ticks must be >= 2")

    network, pre_ids, target_id = _build_convergent_network(config, 0.0)
    if len(weights) != len(pre_ids):
        raise ValueError("weight vector does not match experiment topology")
    for pre_id, weight in zip(pre_ids, weights):
        network.synapses[pre_id][0].weight = float(weight)

    for pre_id in pre_ids:
        network.inject_current(pre_id, drive)

    peak_v = network.neurons[target_id].v
    spike_tick: int | None = None
    for _ in range(probe_ticks):
        result = network.step()
        peak_v = max(peak_v, network.neurons[target_id].v)
        if target_id in result.spike_ids and spike_tick is None:
            spike_tick = result.tick
    return spike_tick is not None, peak_v, spike_tick


def train_learning_weights(
    config: Config, condition: str
) -> tuple[tuple[float, ...], LearningEngine, TrialPartitions]:
    """Public deterministic training boundary for registered research protocols."""
    return _train(config, condition)


def probe_learning_response(
    config: Config, weights: Sequence[float]
) -> tuple[bool, float, int | None]:
    """Public deterministic post-training probe boundary."""
    return _probe_response(config, weights)


def run_learning_experiment(
    config: Config, condition: str = "learning_on"
) -> LearningExperimentResult:
    """Run training and compare fresh baseline/trained network responses."""
    exp = _experiment_config(config)
    initial_weight = float(exp.get("initial_weight", 0.05))
    pre_count = int(exp.get("presynaptic_neurons", 48))
    initial_weights = tuple(initial_weight for _ in range(pre_count))
    partitions = _validated_trial_partitions(config)

    baseline_spiked, baseline_peak_v, baseline_tick = _probe_response(
        config, initial_weights
    )
    trained_weights, learning, partitions = _train(config, condition)
    trained_spiked, trained_peak_v, trained_tick = _probe_response(
        config, trained_weights
    )

    initial_mean = statistics.mean(initial_weights)
    final_mean = statistics.mean(trained_weights)
    return LearningExperimentResult(
        training_trials=int(exp.get("training_trials", 20)),
        presynaptic_neurons=pre_count,
        initial_mean_weight=initial_mean,
        final_mean_weight=final_mean,
        mean_weight_delta=final_mean - initial_mean,
        rewards_received=learning.stats.rewards_received,
        rewards_applied=learning.stats.rewards_applied,
        reward_weight_updates=learning.stats.reward_weight_updates,
        baseline_target_spiked=baseline_spiked,
        trained_target_spiked=trained_spiked,
        baseline_target_peak_v=baseline_peak_v,
        trained_target_peak_v=trained_peak_v,
        baseline_target_spike_tick=baseline_tick,
        trained_target_spike_tick=trained_tick,
        condition=condition,
        train_trial_count=len(partitions["train"]),
        validation_trial_count=len(partitions["validation"]),
        holdout_trial_count=len(partitions["holdout"]),
        protocol_id=str(exp["protocol_id"]),
        protocol_version=int(exp["protocol_version"]),
    )


def _load_yaml(path: Path) -> dict[str, Any]:
    """Load and validate a YAML configuration file."""
    with path.open("r", encoding="utf-8") as handle:
        loaded = yaml.safe_load(handle)
    if not isinstance(loaded, dict):
        raise TypeError("experiment config root must be a mapping")
    return cast(dict[str, Any], loaded)


def main() -> int:
    """CLI entry point for the deterministic learning experiment."""
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--config", default="configs/learning_experiment.yaml")
    args = parser.parse_args()
    result = run_learning_experiment(_load_yaml(Path(args.config)))
    for key, value in asdict(result).items():
        print(f"{key}: {value}")
    print(f"learned: {result.learned}")
    return 0 if result.learned else 1


if __name__ == "__main__":
    raise SystemExit(main())