Spaces:
Paused
Paused
Download src/experiments/learning_lab.py from ThomasHeisig/Brain-5D-Space: direct link, hf CLI and curl.
- Browser
- Download file 15.4 kB
-
https://huggingface.co/spaces/ThomasHeisig/Brain-5D-Space/resolve/main/src/experiments/learning_lab.py
- Command line
-
hf download hf://spaces/ThomasHeisig/Brain-5D-Space/src/experiments/learning_lab.py
-
curl -L -o learning_lab.py https://huggingface.co/spaces/ThomasHeisig/Brain-5D-Space/resolve/main/src/experiments/learning_lab.py
15.4 kB
| """Deterministic end-to-end learning experiment for Brain 5D. | |
| 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, ...]] | |
| 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" | |
| 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()) | |