diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..0f2801be95e7863ee74dc053aed37e230ce75fd2 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,15 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +3D_SEMMN_final.pdf filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/manifold_pca.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/neuron_clusters_3d.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/neuron_module_audio.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/neuron_module_hub.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/neuron_module_other.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/neuron_module_vision.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/neuron_modules_all.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/route_map_audio_2d.png filter=lfs diff=lfs merge=lfs -text +3D-SEMMN/outputs/figures/route_map_vision_2d.png filter=lfs diff=lfs merge=lfs -text diff --git a/3D-SEMMN/.DS_Store b/3D-SEMMN/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..fb29c4f5c6ddc3216305f0963413cd46faae030f Binary files /dev/null and b/3D-SEMMN/.DS_Store differ diff --git a/3D-SEMMN/README.md b/3D-SEMMN/README.md new file mode 100644 index 0000000000000000000000000000000000000000..be6db5256c094e7f136b4a68d16e64400da1788d --- /dev/null +++ b/3D-SEMMN/README.md @@ -0,0 +1,105 @@ +# 3D-SEMMN (Prototype) + +Minimal but fully functional PyTorch 2.3+ prototype of a **3D Spatially-Embedded Multimodal Manifold Network** with: + +- fixed 3D neuron coordinates (`100 x 100 x 50` in `config/default.yaml`) +- 100 columnar modules (dense intra-column, sparse inter-column) +- exponential distance-penalized sparse connectivity (`torch.sparse`) +- Hebbian/STDP update every 100 steps +- shared perceptual manifold hub with cross-modal attention +- manifold regularization (latent `d=8`) + geometric twist layer (Cayley rotation) +- multimodal losses: classification + contrastive + reconstruction + imagination cycle +- ablations for spatial penalty and manifold loss + +This code is designed for readability and modular experimentation rather than full biological fidelity. + +## Biological Inspirations + +- **seRNN (Achterberg et al., 2023)**: spatial embedding + wiring-cost constraints for modular sparse topology. +- **MSeNN-style multimodal ideas (2021)**: cross-modal coupling and biologically motivated integration. +- **Gallego et al. (2017)**: low-dimensional population manifolds. + +## Project Layout + +- `semmn/config.py`: dataclass config + YAML loader +- `semmn/grid.py`: 3D coordinates, lobe masks, column assignment +- `semmn/connectivity.py`: sparse distance-decay edges + STDP + energy proxy +- `semmn/columns.py`: low-rank dense intra-column dynamics +- `semmn/encoders.py`: vision CNN and audio 1D-conv encoder +- `semmn/manifold.py`: manifold autoencoder and geometric twist +- `semmn/hub.py`: shared manifold hub + attention + imagination decoders +- `semmn/losses.py`: full training objective +- `semmn/model.py`: full SEMMN forward pass +- `semmn/dataset.py`: AV-MNIST loading with synthetic fallback +- `train.py`: training/eval/checkpointing +- `visualize.py`: 3D grid, manifold PCA, metrics plots + +## Quick Start + +```bash +python -m venv .venv +source .venv/bin/activate +pip install -r requirements.txt +``` + +Train (light config, recommended first): + +```bash +python train.py --config config/light.yaml --device cuda +``` + +Train (full 500K-neuron spatial grid metadata): + +```bash +python train.py --config config/default.yaml --device cuda +``` + +Generate figures: + +```bash +python visualize.py \ + --config config/light.yaml \ + --checkpoint outputs/best.pt \ + --history outputs/history.json \ + --device cuda +``` + +## Config Notes + +- `config/default.yaml`: full spatial grid (`100x100x50`) +- `config/light.yaml`: faster profile (`100x100x10`) for rapid iteration +- Ablations: + - `loss.use_spatial_penalty: false` + - `loss.use_manifold_loss: false` + +## Reported Metrics + +`train.py` logs: + +- classification loss and total loss +- validation classification accuracy +- cross-modal retrieval accuracy (vision->audio nearest match in batch) +- active sparse parameters (`nnz` proxy) +- energy proxy (`mean(|w| * distance)`) + +## Extension Notes + +### Full spiking variant (Brian2 / snnTorch) + +1. Replace rate-based state updates (`tanh`) with LIF/AdEx neuron dynamics. +2. Convert STDP update from rate Hebbian correlation to spike-timing windows. +3. Keep 3D coordinates and sparse wiring graph unchanged for fair ablations. + +### Neuromorphic export direction + +1. Map columns to compute cores (Loihi/SpiNNaker style partitioning). +2. Convert sparse inter-column edges to routing tables. +3. Quantize state variables and weights (8-bit or mixed precision). +4. Use event-driven updates for power-efficient deployment. + +## Practical Runtime Guidance + +- Start with `config/light.yaml` to confirm setup and debug quickly. +- Increase `training.batch_size` and `model.recurrent_steps` gradually. +- Keep STDP interval at `100` to limit plasticity overhead. +- Use mixed precision (`training.amp: true`) on CUDA devices. diff --git a/3D-SEMMN/__pycache__/train.cpython-312.pyc b/3D-SEMMN/__pycache__/train.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..0c9af2df0eacf9a2efc1e9a3dae8783b276fe1fa Binary files /dev/null and b/3D-SEMMN/__pycache__/train.cpython-312.pyc differ diff --git a/3D-SEMMN/__pycache__/visualize.cpython-312.pyc b/3D-SEMMN/__pycache__/visualize.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..9a99fab9ee9b92eda1644f3f1ca7b9ef50abee50 Binary files /dev/null and b/3D-SEMMN/__pycache__/visualize.cpython-312.pyc differ diff --git a/3D-SEMMN/config/default.yaml b/3D-SEMMN/config/default.yaml new file mode 100644 index 0000000000000000000000000000000000000000..e6188051924b7c09efad74a8bbee707a9c76390e --- /dev/null +++ b/3D-SEMMN/config/default.yaml @@ -0,0 +1,51 @@ +seed: 42 +grid: + dims: [100, 100, 50] + n_columns_xy: [10, 10] + vision_y_ratio: 0.30 + audio_y_ratio: 0.70 + hub_y_range: [0.40, 0.60] + hub_z_range: [0.30, 0.70] + +model: + state_dim: 32 + latent_dim: 8 + vision_feature_dim: 256 + audio_feature_dim: 256 + recurrent_steps: 6 + attention_heads: 4 + +connectivity: + lambda_decay: 5.0 + inter_density: 0.08 + intra_density_boost: 2.0 + min_weight: 0.0 + max_weight: 1.5 + init_weight_scale: 0.15 + stdp_interval: 100 + stdp_lr: 0.01 + anti_hebbian: 0.05 + +loss: + contrastive_weight: 1.0 + reconstruction_weight: 0.5 + imagination_weight: 0.3 + manifold_weight: 0.1 + spatial_penalty_weight: 0.0005 + temperature: 0.1 + use_spatial_penalty: true + use_manifold_loss: true + +training: + dataset_source: synthetic + batch_size: 64 + epochs: 10 + lr: 0.001 + weight_decay: 0.0001 + num_workers: 4 + log_interval: 50 + val_batches: 100 + amp: true + +paths: + output_dir: outputs diff --git a/3D-SEMMN/config/light.yaml b/3D-SEMMN/config/light.yaml new file mode 100644 index 0000000000000000000000000000000000000000..7d9d016f8418755a56faab4cabfc74b537792bc4 --- /dev/null +++ b/3D-SEMMN/config/light.yaml @@ -0,0 +1,17 @@ +seed: 42 +grid: + dims: [100, 100, 10] + n_columns_xy: [10, 10] + +model: + recurrent_steps: 4 + state_dim: 24 + +connectivity: + inter_density: 0.10 + +training: + batch_size: 128 + epochs: 10 + num_workers: 2 + log_interval: 20 diff --git a/3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte b/3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte new file mode 100644 index 0000000000000000000000000000000000000000..d026debec174b65df0cd4d448668d0c744497faa --- /dev/null +++ b/3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0fa7898d509279e482958e8ce81c8e77db3f2f8254e26661ceb7762c4d494ce7 +size 7840016 diff --git a/3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte.gz b/3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte.gz new file mode 100644 index 0000000000000000000000000000000000000000..aa17dfe485689242a90be276702dcadd17d406f4 --- /dev/null +++ b/3D-SEMMN/data/MNIST/raw/t10k-images-idx3-ubyte.gz @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d422c7b0a1c1c79245a5bcf07fe86e33eeafee792b84584aec276f5a2dbc4e6 +size 1648877 diff --git a/3D-SEMMN/data/MNIST/raw/t10k-labels-idx1-ubyte b/3D-SEMMN/data/MNIST/raw/t10k-labels-idx1-ubyte new file mode 100644 index 0000000000000000000000000000000000000000..d1c3a970612bbd2df47a3c0697f82bd394abc450 Binary files /dev/null and b/3D-SEMMN/data/MNIST/raw/t10k-labels-idx1-ubyte differ diff --git a/3D-SEMMN/data/MNIST/raw/t10k-labels-idx1-ubyte.gz b/3D-SEMMN/data/MNIST/raw/t10k-labels-idx1-ubyte.gz new file mode 100644 index 0000000000000000000000000000000000000000..d1995bebe8e5b3faeaae99149ce4eb7a68c5764d --- /dev/null +++ b/3D-SEMMN/data/MNIST/raw/t10k-labels-idx1-ubyte.gz @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f7ae60f92e00ec6debd23a6088c31dbd2371eca3ffa0defaefb259924204aec6 +size 4542 diff --git a/3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte b/3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte new file mode 100644 index 0000000000000000000000000000000000000000..6f2c0cdbab4d9d6c202cefd16ea17c10b9e783e2 --- /dev/null +++ b/3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba891046e6505d7aadcbbe25680a0738ad16aec93bde7f9b65e87a2fc25776db +size 47040016 diff --git a/3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte.gz b/3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte.gz new file mode 100644 index 0000000000000000000000000000000000000000..9e9852c14333d6b633709fec2c6df84941243c9d --- /dev/null +++ b/3D-SEMMN/data/MNIST/raw/train-images-idx3-ubyte.gz @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:440fcabf73cc546fa21475e81ea370265605f56be210a4024d2ca8f203523609 +size 9912422 diff --git a/3D-SEMMN/data/MNIST/raw/train-labels-idx1-ubyte b/3D-SEMMN/data/MNIST/raw/train-labels-idx1-ubyte new file mode 100644 index 0000000000000000000000000000000000000000..d6b4c5db3b52063d543fb397aede09aba0dc5234 Binary files /dev/null and b/3D-SEMMN/data/MNIST/raw/train-labels-idx1-ubyte differ diff --git a/3D-SEMMN/data/MNIST/raw/train-labels-idx1-ubyte.gz b/3D-SEMMN/data/MNIST/raw/train-labels-idx1-ubyte.gz new file mode 100644 index 0000000000000000000000000000000000000000..a7ebf9b5b685e9014530844158807071ae717f7f --- /dev/null +++ b/3D-SEMMN/data/MNIST/raw/train-labels-idx1-ubyte.gz @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3552534a0a558bbed6aed32b30c495cca23d567ec52cac8be1a0730e8010255c +size 28881 diff --git a/3D-SEMMN/outputs/.DS_Store b/3D-SEMMN/outputs/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..ffa081ccca8cc6f40d6ed49335c6d95caf701a2a Binary files /dev/null and b/3D-SEMMN/outputs/.DS_Store differ diff --git a/3D-SEMMN/outputs/best.pt b/3D-SEMMN/outputs/best.pt new file mode 100644 index 0000000000000000000000000000000000000000..b4e4fb50bcfa52d57c84d5d64cc1dc95b456d1b4 --- /dev/null +++ b/3D-SEMMN/outputs/best.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:66d1a8bbb97532ad3ea4005ad99e6dbda5b5cced507c3d5c6961fb52b0b1b1a2 +size 7554761 diff --git a/3D-SEMMN/outputs/figures/manifold_pca.png b/3D-SEMMN/outputs/figures/manifold_pca.png new file mode 100644 index 0000000000000000000000000000000000000000..427bae4400a16619b366e11ac30ce7de4ee30d52 --- /dev/null +++ b/3D-SEMMN/outputs/figures/manifold_pca.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cac438e82d7ca1cf2f52af7708f6b9335a8663479adcb74ee24df4a83944a50b +size 118297 diff --git a/3D-SEMMN/outputs/figures/metrics.png b/3D-SEMMN/outputs/figures/metrics.png new file mode 100644 index 0000000000000000000000000000000000000000..4354c96376b3c1bab367dde015cfcef711380a9e Binary files /dev/null and b/3D-SEMMN/outputs/figures/metrics.png differ diff --git a/3D-SEMMN/outputs/figures/neuron_clusters_3d.png b/3D-SEMMN/outputs/figures/neuron_clusters_3d.png new file mode 100644 index 0000000000000000000000000000000000000000..023ebd2c679021a40893a242849c87533dacd5d5 --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_clusters_3d.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ed62b63fe58c290a70037de84ad4f7b1b4c8f239d8028ce374f85477dac6affc +size 1416937 diff --git a/3D-SEMMN/outputs/figures/neuron_grid_interactive.html b/3D-SEMMN/outputs/figures/neuron_grid_interactive.html new file mode 100644 index 0000000000000000000000000000000000000000..1c6a58eda9b88335d85f87e7122ca70ed5760d0b --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_grid_interactive.html @@ -0,0 +1,7 @@ + + + +
+
+ + \ No newline at end of file diff --git a/3D-SEMMN/outputs/figures/neuron_module_audio.png b/3D-SEMMN/outputs/figures/neuron_module_audio.png new file mode 100644 index 0000000000000000000000000000000000000000..1a3cb423fda22ab9bc007a3eb8d81e9fa8e2d8f8 --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_module_audio.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7308eb43137af236f224879b820d2ad4f4a590183f2638c8831e30a8e5751e1e +size 1220990 diff --git a/3D-SEMMN/outputs/figures/neuron_module_hub.png b/3D-SEMMN/outputs/figures/neuron_module_hub.png new file mode 100644 index 0000000000000000000000000000000000000000..fef4615fbedd510b564f9031825df161f0d9bd1e --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_module_hub.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5240e02ec0f7f6a58e6e92131a12a049696ba8f4e6a5fcbb3fd1c0ff7122f5a9 +size 938952 diff --git a/3D-SEMMN/outputs/figures/neuron_module_other.png b/3D-SEMMN/outputs/figures/neuron_module_other.png new file mode 100644 index 0000000000000000000000000000000000000000..41a5af0b98a268f259b24b36710345c1244b120a --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_module_other.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63ec4360740e7f10fa93ef706ff71d719540624bcbbdd0690b473c0b04bc908c +size 1212309 diff --git a/3D-SEMMN/outputs/figures/neuron_module_vision.png b/3D-SEMMN/outputs/figures/neuron_module_vision.png new file mode 100644 index 0000000000000000000000000000000000000000..0d4621772e71839b69594b0277f6c5fccab2c115 --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_module_vision.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0b2041bcf4e54f8b43f4e42ee2f1991c0e82a0e731cc54a5d76f0c58f60facf +size 1232056 diff --git a/3D-SEMMN/outputs/figures/neuron_modules_all.png b/3D-SEMMN/outputs/figures/neuron_modules_all.png new file mode 100644 index 0000000000000000000000000000000000000000..7e6c2718995b1e824bb4c7b61e66451506d08f40 --- /dev/null +++ b/3D-SEMMN/outputs/figures/neuron_modules_all.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:98172e67fae1ad85784bf5ed071e9189946dcc665a28508e947b5a59c4415ca8 +size 1801792 diff --git a/3D-SEMMN/outputs/figures/route_edges.csv b/3D-SEMMN/outputs/figures/route_edges.csv new file mode 100644 index 0000000000000000000000000000000000000000..eadfe9e61797dc7c1dc58237e5ebf0204d21b47d --- /dev/null +++ b/3D-SEMMN/outputs/figures/route_edges.csv @@ -0,0 +1,13 @@ +step,source,target,strength,modality +0.0,80.0,90.0,0.014531210064888,vision +0.0,10.0,0.0,0.0044663818553090096,vision +0.0,22.0,11.0,0.003179906401783228,vision +1.0,80.0,90.0,0.013347183354198933,vision +1.0,10.0,0.0,0.004102454055100679,vision +1.0,22.0,11.0,0.002920802216976881,vision +2.0,80.0,90.0,0.012341569177806377,vision +2.0,10.0,0.0,0.0037933634594082832,vision +2.0,22.0,11.0,0.0027007409371435642,vision +3.0,80.0,90.0,0.011476870626211166,vision +3.0,10.0,0.0,0.0035275856498628855,vision +3.0,22.0,11.0,0.0025115162134170532,vision diff --git a/3D-SEMMN/outputs/figures/route_map_audio_2d.png b/3D-SEMMN/outputs/figures/route_map_audio_2d.png new file mode 100644 index 0000000000000000000000000000000000000000..9661766d99518bf3f335ae8c29db24ca4d6bd97f --- /dev/null +++ b/3D-SEMMN/outputs/figures/route_map_audio_2d.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:44696dc9efe5a2209bea6e43bdfb463b229acd0ded4a869b0351626791ae4fd0 +size 112247 diff --git a/3D-SEMMN/outputs/figures/route_map_interactive.html b/3D-SEMMN/outputs/figures/route_map_interactive.html new file mode 100644 index 0000000000000000000000000000000000000000..5d219195080d20565a1ee6e3b904e594ea5a1861 --- /dev/null +++ b/3D-SEMMN/outputs/figures/route_map_interactive.html @@ -0,0 +1,7 @@ + + + +
+
+ + \ No newline at end of file diff --git a/3D-SEMMN/outputs/figures/route_map_vision_2d.png b/3D-SEMMN/outputs/figures/route_map_vision_2d.png new file mode 100644 index 0000000000000000000000000000000000000000..c4c145fa4911472bc1253584083abb38847ef14f --- /dev/null +++ b/3D-SEMMN/outputs/figures/route_map_vision_2d.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9836e9a5370eb3bd45986e10bd2868effe1def45ff34aa6e2feaca7ed80e0543 +size 117622 diff --git a/3D-SEMMN/outputs/figures/route_nodes.csv b/3D-SEMMN/outputs/figures/route_nodes.csv new file mode 100644 index 0000000000000000000000000000000000000000..c1dcf6ad257b79f6b229ed6b3ce2d4953f8c9dd3 --- /dev/null +++ b/3D-SEMMN/outputs/figures/route_nodes.csv @@ -0,0 +1,241 @@ +step,column,activity,modality +0.0,0.0,0.02956058643758297,vision +0.0,1.0,0.0409482941031456,vision +0.0,2.0,0.026564355939626694,vision +0.0,10.0,0.040411435067653656,vision +0.0,11.0,0.02947550266981125,vision +0.0,12.0,0.0309763066470623,vision +0.0,20.0,0.027739018201828003,vision +0.0,21.0,0.023138199001550674,vision +0.0,22.0,0.03138161823153496,vision +0.0,30.0,0.03375103697180748,vision +0.0,31.0,0.034890495240688324,vision +0.0,32.0,0.025857770815491676,vision +0.0,40.0,0.03491085767745972,vision +0.0,41.0,0.04038066044449806,vision +0.0,42.0,0.030116962268948555,vision +0.0,50.0,0.029867637902498245,vision +0.0,51.0,0.036975305527448654,vision +0.0,52.0,0.04041629657149315,vision +0.0,60.0,0.03777403011918068,vision +0.0,61.0,0.02775333821773529,vision +0.0,62.0,0.03777413070201874,vision +0.0,70.0,0.029612775892019272,vision +0.0,71.0,0.03889508917927742,vision +0.0,72.0,0.03445589914917946,vision +0.0,80.0,0.03937682509422302,vision +0.0,81.0,0.027864739298820496,vision +0.0,82.0,0.027492109686136246,vision +0.0,90.0,0.036627497524023056,vision +0.0,91.0,0.03462231159210205,vision +0.0,92.0,0.04038887843489647,vision +1.0,0.0,0.04356175288558006,vision +1.0,1.0,0.03761175647377968,vision +1.0,2.0,0.02439984865486622,vision +1.0,10.0,0.03711864724755287,vision +1.0,11.0,0.038757000118494034,vision +1.0,12.0,0.028452306985855103,vision +1.0,20.0,0.025478797033429146,vision +1.0,21.0,0.021252859383821487,vision +1.0,22.0,0.028824592009186745,vision +1.0,30.0,0.03100094571709633,vision +1.0,31.0,0.03204755857586861,vision +1.0,32.0,0.02375083602964878,vision +1.0,40.0,0.032066263258457184,vision +1.0,41.0,0.037090376019477844,vision +1.0,42.0,0.027662981301546097,vision +1.0,50.0,0.02743397280573845,vision +1.0,51.0,0.033962495625019073,vision +1.0,52.0,0.03712311014533043,vision +1.0,60.0,0.03469613939523697,vision +1.0,61.0,0.025491949170827866,vision +1.0,62.0,0.03469623252749443,vision +1.0,70.0,0.0271998792886734,vision +1.0,71.0,0.035725854337215424,vision +1.0,72.0,0.03164837509393692,vision +1.0,80.0,0.03616833686828613,vision +1.0,81.0,0.025594275444746017,vision +1.0,82.0,0.025252006947994232,vision +1.0,90.0,0.0870317593216896,vision +1.0,91.0,0.03180122748017311,vision +1.0,92.0,0.03709792718291283,vision +2.0,0.0,0.055453140288591385,vision +2.0,1.0,0.03477798402309418,vision +2.0,2.0,0.02256149612367153,vision +2.0,10.0,0.03432202339172363,vision +2.0,11.0,0.04663990065455437,vision +2.0,12.0,0.026308631524443626,vision +2.0,20.0,0.023559153079986572,vision +2.0,21.0,0.019651610404253006,vision +2.0,22.0,0.026652866974473,vision +2.0,30.0,0.028665248304605484,vision +2.0,31.0,0.02963300608098507,vision +2.0,32.0,0.021961383521556854,vision +2.0,40.0,0.02965030074119568,vision +2.0,41.0,0.034295883029699326,vision +2.0,42.0,0.025578774511814117,vision +2.0,50.0,0.025367021560668945,vision +2.0,51.0,0.031403668224811554,vision +2.0,52.0,0.03432615101337433,vision +2.0,60.0,0.03208203613758087,vision +2.0,61.0,0.0235713142901659,vision +2.0,62.0,0.032082121819257736,vision +2.0,70.0,0.025150563567876816,vision +2.0,71.0,0.03303416818380356,vision +2.0,72.0,0.029263898730278015,vision +2.0,80.0,0.03344331309199333,vision +2.0,81.0,0.023665931075811386,vision +2.0,82.0,0.023349450901150703,vision +2.0,90.0,0.12984082102775574,vision +2.0,91.0,0.02940523438155651,vision +2.0,92.0,0.03430286794900894,vision +3.0,0.0,0.06567821651697159,vision +3.0,1.0,0.03234130144119263,vision +3.0,2.0,0.020980749279260635,vision +3.0,10.0,0.03191728889942169,vision +3.0,11.0,0.05341819301247597,vision +3.0,12.0,0.024465344846248627,vision +3.0,20.0,0.021908504888415337,vision +3.0,21.0,0.018274743109941483,vision +3.0,22.0,0.024785460904240608,vision +3.0,30.0,0.026656849309802055,vision +3.0,31.0,0.027556801214814186,vision +3.0,32.0,0.020422684028744698,vision +3.0,40.0,0.027572883293032646,vision +3.0,41.0,0.03189297765493393,vision +3.0,42.0,0.023786624893546104,vision +3.0,50.0,0.023589707911014557,vision +3.0,51.0,0.029203403741121292,vision +3.0,52.0,0.03192112594842911,vision +3.0,60.0,0.029834242537617683,vision +3.0,61.0,0.021919816732406616,vision +3.0,62.0,0.0298343226313591,vision +3.0,70.0,0.023388417437672615,vision +3.0,71.0,0.030719663947820663,vision +3.0,72.0,0.027213554829359055,vision +3.0,80.0,0.031100144609808922,vision +3.0,81.0,0.02200780250132084,vision +3.0,82.0,0.02171349711716175,vision +3.0,90.0,0.1666511446237564,vision +3.0,91.0,0.02734498865902424,vision +3.0,92.0,0.03189947456121445,vision +0.0,7.0,0.03473106026649475,audio +0.0,8.0,0.02794949896633625,audio +0.0,9.0,0.027408087626099586,audio +0.0,17.0,0.038603682070970535,audio +0.0,18.0,0.036946672946214676,audio +0.0,19.0,0.02548990212380886,audio +0.0,27.0,0.03835974261164665,audio +0.0,28.0,0.04123900458216667,audio +0.0,29.0,0.030051397159695625,audio +0.0,37.0,0.03510997071862221,audio +0.0,38.0,0.03655286878347397,audio +0.0,39.0,0.024867547675967216,audio +0.0,47.0,0.039415426552295685,audio +0.0,48.0,0.0356949046254158,audio +0.0,49.0,0.04190101474523544,audio +0.0,57.0,0.033966340124607086,audio +0.0,58.0,0.03714350610971451,audio +0.0,59.0,0.03321348875761032,audio +0.0,67.0,0.030020426958799362,audio +0.0,68.0,0.034492406994104385,audio +0.0,69.0,0.02616014890372753,audio +0.0,77.0,0.03146708384156227,audio +0.0,78.0,0.04283866658806801,audio +0.0,79.0,0.02919699251651764,audio +0.0,87.0,0.03697360306978226,audio +0.0,88.0,0.03308973088860512,audio +0.0,89.0,0.02546405978500843,audio +0.0,97.0,0.03400813415646553,audio +0.0,98.0,0.02903173305094242,audio +0.0,99.0,0.028612902387976646,audio +1.0,7.0,0.03473105654120445,audio +1.0,8.0,0.0279494971036911,audio +1.0,9.0,0.027408085763454437,audio +1.0,17.0,0.03860367834568024,audio +1.0,18.0,0.03694666922092438,audio +1.0,19.0,0.02548990026116371,audio +1.0,27.0,0.038359738886356354,audio +1.0,28.0,0.04123900085687637,audio +1.0,29.0,0.030051395297050476,audio +1.0,37.0,0.03510996699333191,audio +1.0,38.0,0.03655286505818367,audio +1.0,39.0,0.024867547675967216,audio +1.0,47.0,0.039415426552295685,audio +1.0,48.0,0.035694900900125504,audio +1.0,49.0,0.04190101474523544,audio +1.0,57.0,0.03396633639931679,audio +1.0,58.0,0.03714350238442421,audio +1.0,59.0,0.03321348503232002,audio +1.0,67.0,0.030020423233509064,audio +1.0,68.0,0.03449240326881409,audio +1.0,69.0,0.026160147041082382,audio +1.0,77.0,0.03146708011627197,audio +1.0,78.0,0.04283866286277771,audio +1.0,79.0,0.02919699065387249,audio +1.0,87.0,0.03697359934449196,audio +1.0,88.0,0.03308972716331482,audio +1.0,89.0,0.02546405978500843,audio +1.0,97.0,0.03400813415646553,audio +1.0,98.0,0.02903173118829727,audio +1.0,99.0,0.028612900525331497,audio +2.0,7.0,0.03473105654120445,audio +2.0,8.0,0.02794949896633625,audio +2.0,9.0,0.027408087626099586,audio +2.0,17.0,0.038603682070970535,audio +2.0,18.0,0.036946672946214676,audio +2.0,19.0,0.02548990212380886,audio +2.0,27.0,0.03835974261164665,audio +2.0,28.0,0.04123900458216667,audio +2.0,29.0,0.030051397159695625,audio +2.0,37.0,0.03510997071862221,audio +2.0,38.0,0.03655286878347397,audio +2.0,39.0,0.024867551401257515,audio +2.0,47.0,0.03941543027758598,audio +2.0,48.0,0.035694900900125504,audio +2.0,49.0,0.04190101847052574,audio +2.0,57.0,0.033966340124607086,audio +2.0,58.0,0.03714350610971451,audio +2.0,59.0,0.03321348503232002,audio +2.0,67.0,0.030020425096154213,audio +2.0,68.0,0.034492406994104385,audio +2.0,69.0,0.02616014890372753,audio +2.0,77.0,0.03146708384156227,audio +2.0,78.0,0.04283866658806801,audio +2.0,79.0,0.02919699251651764,audio +2.0,87.0,0.03697360306978226,audio +2.0,88.0,0.03308973088860512,audio +2.0,89.0,0.02546406351029873,audio +2.0,97.0,0.03400813788175583,audio +2.0,98.0,0.02903173305094242,audio +2.0,99.0,0.028612902387976646,audio +3.0,7.0,0.03473105654120445,audio +3.0,8.0,0.0279495008289814,audio +3.0,9.0,0.027408087626099586,audio +3.0,17.0,0.038603682070970535,audio +3.0,18.0,0.036946672946214676,audio +3.0,19.0,0.02548990212380886,audio +3.0,27.0,0.03835974261164665,audio +3.0,28.0,0.04123900458216667,audio +3.0,29.0,0.030051397159695625,audio +3.0,37.0,0.03510997071862221,audio +3.0,38.0,0.03655286878347397,audio +3.0,39.0,0.024867551401257515,audio +3.0,47.0,0.03941543027758598,audio +3.0,48.0,0.035694900900125504,audio +3.0,49.0,0.04190101847052574,audio +3.0,57.0,0.033966340124607086,audio +3.0,58.0,0.03714350610971451,audio +3.0,59.0,0.03321348503232002,audio +3.0,67.0,0.030020426958799362,audio +3.0,68.0,0.034492406994104385,audio +3.0,69.0,0.02616015076637268,audio +3.0,77.0,0.03146708384156227,audio +3.0,78.0,0.04283866658806801,audio +3.0,79.0,0.02919699251651764,audio +3.0,87.0,0.03697360306978226,audio +3.0,88.0,0.03308973088860512,audio +3.0,89.0,0.02546406351029873,audio +3.0,97.0,0.03400813788175583,audio +3.0,98.0,0.02903173491358757,audio +3.0,99.0,0.028612902387976646,audio diff --git a/3D-SEMMN/outputs/history.json b/3D-SEMMN/outputs/history.json new file mode 100644 index 0000000000000000000000000000000000000000..47605c300ae0697996fccbf28c6f550ba9ed3bfd --- /dev/null +++ b/3D-SEMMN/outputs/history.json @@ -0,0 +1,37 @@ +[ + { + "epoch": 0, + "step": 0, + "loss_total": 6.665867805480957, + "loss_cls": 2.322234630584717, + "loss_contrastive": 4.330072402954102, + "loss_manifold": 0.007405777927488089, + "active_params": 2.0, + "energy_proxy": 0.17007668316364288 + }, + { + "epoch": 0, + "step": 10, + "loss_total": 5.909451484680176, + "loss_cls": 2.2938742637634277, + "loss_contrastive": 3.5930726528167725, + "loss_manifold": 0.001445187022909522, + "active_params": 6.0, + "energy_proxy": 0.12415368109941483 + }, + { + "epoch": 0, + "step": 20, + "loss_total": 5.335961818695068, + "loss_cls": 2.2907979488372803, + "loss_contrastive": 3.0338499546051025, + "loss_manifold": 0.0009116176515817642, + "active_params": 6.0, + "energy_proxy": 0.07187842577695847 + }, + { + "epoch": 0, + "val_accuracy": 0.0, + "val_cross_modal_retrieval": 0.0 + } +] \ No newline at end of file diff --git a/3D-SEMMN/requirements.txt b/3D-SEMMN/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..f48b8ba08417c34307a472594a1c64068b880053 --- /dev/null +++ b/3D-SEMMN/requirements.txt @@ -0,0 +1,10 @@ +torch>=2.3 +torchvision>=0.18 +numpy>=1.25 +pandas>=2.2 +PyYAML>=6.0 +matplotlib>=3.8 +plotly>=5.22 +scikit-learn>=1.4 +tqdm>=4.66 +datasets>=2.20 diff --git a/3D-SEMMN/semmn/__init__.py b/3D-SEMMN/semmn/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..bdcde8efc3913745644965067afade130a98cf19 --- /dev/null +++ b/3D-SEMMN/semmn/__init__.py @@ -0,0 +1,6 @@ +"""3D-SEMMN package.""" + +from semmn.config import SEMMNConfig +from semmn.model import SEMMNModel + +__all__ = ["SEMMNConfig", "SEMMNModel"] diff --git a/3D-SEMMN/semmn/__pycache__/__init__.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..634130048925ef6e41d153286ba14e9af74d4705 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/__init__.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/columns.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/columns.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..f55986a200f106ca17ab352fb97bb49f2680f5e3 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/columns.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/config.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/config.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..76210f0372341ff5f7e2ca9bb5b18682156ea95b Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/config.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/connectivity.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/connectivity.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..258416d0684373c2320de97f97b7f157fcdc1ff4 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/connectivity.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/dataset.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/dataset.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..8573f397571ef034533218da8a7a8009d927be4c Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/dataset.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/encoders.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/encoders.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..d4999081e9b2fb9c2606e3bd08b20d5b44c18870 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/encoders.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/grid.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/grid.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..3c45785bed42e9746dadc2c37de4d23ed11f1017 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/grid.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/hub.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/hub.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..b24385959861e40e1170edf667c03e2d38570435 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/hub.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/losses.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/losses.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..51b906c33556c58533abf098cc0ac6f0ce652096 Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/losses.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/manifold.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/manifold.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a73cda5bc55318e775f0b9a0477ed5b04f0541cb Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/manifold.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/__pycache__/model.cpython-312.pyc b/3D-SEMMN/semmn/__pycache__/model.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..d60c30c7772cf599099a70c7c6084c999ea43dbe Binary files /dev/null and b/3D-SEMMN/semmn/__pycache__/model.cpython-312.pyc differ diff --git a/3D-SEMMN/semmn/columns.py b/3D-SEMMN/semmn/columns.py new file mode 100644 index 0000000000000000000000000000000000000000..21b9952696f5d64ee7f25eb5f3eed49b89cd9e47 --- /dev/null +++ b/3D-SEMMN/semmn/columns.py @@ -0,0 +1,33 @@ +"""Columnar dynamics modules.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + + +class ColumnarDynamics(nn.Module): + """Dense intra-column dynamics with low-rank factors.""" + + def __init__(self, num_columns: int, state_dim: int, rank: int = 8) -> None: + super().__init__() + self.num_columns = num_columns + self.state_dim = state_dim + self.rank = rank + + self.u = nn.Parameter(torch.randn(num_columns, state_dim, rank) * 0.05) + self.v = nn.Parameter(torch.randn(num_columns, rank, state_dim) * 0.05) + self.bias = nn.Parameter(torch.zeros(num_columns, state_dim)) + self.norm = nn.LayerNorm(state_dim) + + def forward(self, states: torch.Tensor, inter_messages: torch.Tensor) -> torch.Tensor: + # W_col = U @ V in factored form to keep parameters compact. + w = torch.matmul(self.u, self.v) # (C, D, D) + intra = torch.einsum("bcd,cde->bce", states, w) + updated = intra + inter_messages + self.bias.unsqueeze(0) + return torch.tanh(self.norm(updated)) + + +def pool_columns(states: torch.Tensor) -> torch.Tensor: + """Convert `(B, C, D)` into `(B, C)` pooled activity.""" + return states.mean(dim=-1) diff --git a/3D-SEMMN/semmn/config.py b/3D-SEMMN/semmn/config.py new file mode 100644 index 0000000000000000000000000000000000000000..ed18810746596fa42f7dd177e9b6a80965daab64 --- /dev/null +++ b/3D-SEMMN/semmn/config.py @@ -0,0 +1,121 @@ +"""Configuration handling for 3D-SEMMN.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import yaml + + +@dataclass +class GridConfig: + dims: tuple[int, int, int] = (100, 100, 50) + n_columns_xy: tuple[int, int] = (10, 10) + vision_y_ratio: float = 0.30 + audio_y_ratio: float = 0.70 + hub_y_range: tuple[float, float] = (0.40, 0.60) + hub_z_range: tuple[float, float] = (0.30, 0.70) + + +@dataclass +class ModelConfig: + state_dim: int = 32 + latent_dim: int = 8 + vision_feature_dim: int = 256 + audio_feature_dim: int = 256 + recurrent_steps: int = 6 + attention_heads: int = 4 + + +@dataclass +class ConnectivityConfig: + lambda_decay: float = 5.0 + inter_density: float = 0.08 + intra_density_boost: float = 2.0 + min_weight: float = 0.0 + max_weight: float = 1.5 + init_weight_scale: float = 0.15 + stdp_interval: int = 100 + stdp_lr: float = 0.01 + anti_hebbian: float = 0.05 + + +@dataclass +class LossConfig: + contrastive_weight: float = 1.0 + reconstruction_weight: float = 0.5 + imagination_weight: float = 0.3 + manifold_weight: float = 0.1 + spatial_penalty_weight: float = 5e-4 + temperature: float = 0.1 + use_spatial_penalty: bool = True + use_manifold_loss: bool = True + + +@dataclass +class TrainingConfig: + dataset_source: str = "synthetic" + batch_size: int = 64 + epochs: int = 10 + lr: float = 1e-3 + weight_decay: float = 1e-4 + num_workers: int = 4 + log_interval: int = 50 + val_batches: int = 100 + amp: bool = True + + +@dataclass +class PathsConfig: + output_dir: str = "outputs" + + +@dataclass +class SEMMNConfig: + seed: int = 42 + grid: GridConfig = field(default_factory=GridConfig) + model: ModelConfig = field(default_factory=ModelConfig) + connectivity: ConnectivityConfig = field(default_factory=ConnectivityConfig) + loss: LossConfig = field(default_factory=LossConfig) + training: TrainingConfig = field(default_factory=TrainingConfig) + paths: PathsConfig = field(default_factory=PathsConfig) + + @property + def n_columns(self) -> int: + return self.grid.n_columns_xy[0] * self.grid.n_columns_xy[1] + + @property + def n_neurons(self) -> int: + return self.grid.dims[0] * self.grid.dims[1] * self.grid.dims[2] + + @classmethod + def from_yaml(cls, path: str | Path) -> "SEMMNConfig": + with Path(path).open("r", encoding="utf-8") as f: + raw = yaml.safe_load(f) or {} + return cls.from_dict(raw) + + @classmethod + def from_dict(cls, raw: dict[str, Any]) -> "SEMMNConfig": + return cls( + seed=raw.get("seed", 42), + grid=_merge_dataclass(GridConfig(), raw.get("grid", {})), + model=_merge_dataclass(ModelConfig(), raw.get("model", {})), + connectivity=_merge_dataclass( + ConnectivityConfig(), raw.get("connectivity", {}) + ), + loss=_merge_dataclass(LossConfig(), raw.get("loss", {})), + training=_merge_dataclass(TrainingConfig(), raw.get("training", {})), + paths=_merge_dataclass(PathsConfig(), raw.get("paths", {})), + ) + + +def _merge_dataclass(default_obj: Any, overrides: dict[str, Any]) -> Any: + current = default_obj.__dict__.copy() + for key, value in overrides.items(): + if isinstance(current.get(key), tuple) and isinstance(value, list): + current[key] = tuple(value) + else: + current[key] = value + return type(default_obj)(**current) diff --git a/3D-SEMMN/semmn/connectivity.py b/3D-SEMMN/semmn/connectivity.py new file mode 100644 index 0000000000000000000000000000000000000000..660a495bf904908b4c47893ab964b7904d05c988 --- /dev/null +++ b/3D-SEMMN/semmn/connectivity.py @@ -0,0 +1,98 @@ +"""Sparse distance-penalized connectivity and STDP rules.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + +from semmn.config import ConnectivityConfig + + +class SparseInterColumnConnectivity(nn.Module): + """Sparse inter-column graph with distance-decay sampling. + + Inspired by seRNN wiring constraints: long-range links carry a larger cost, + so we initialize and regularize them with an exponential distance penalty. + """ + + def __init__( + self, + column_centers: torch.Tensor, + state_dim: int, + config: ConnectivityConfig, + ) -> None: + super().__init__() + self.state_dim = state_dim + self.config = config + self.num_columns = int(column_centers.shape[0]) + self.register_buffer("column_centers", column_centers.clone()) + + row_idx, col_idx, distances, weights = self._build_sparse_edges() + self.register_buffer("row_idx", row_idx) + self.register_buffer("col_idx", col_idx) + self.register_buffer("edge_distances", distances) + self.edge_weights = nn.Parameter(weights) + self.inter_gain = nn.Parameter(torch.ones(state_dim)) + + def _build_sparse_edges(self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + centers = self.column_centers + c = centers.shape[0] + dist = torch.cdist(centers, centers, p=2) + decay = torch.exp(-dist / max(self.config.lambda_decay, 1e-6)) + eye_mask = ~torch.eye(c, dtype=torch.bool, device=centers.device) + prob = (decay * self.config.inter_density).clamp(0.0, 1.0) + sampled = (torch.rand_like(prob) < prob) & eye_mask + row_idx, col_idx = torch.where(sampled) + + if row_idx.numel() == 0: + row_idx = torch.arange(0, c - 1, device=centers.device) + col_idx = torch.arange(1, c, device=centers.device) + + distances = dist[row_idx, col_idx] + weights = torch.randn(row_idx.numel(), device=centers.device) * self.config.init_weight_scale + weights = weights.clamp(self.config.min_weight, self.config.max_weight) + return row_idx.long(), col_idx.long(), distances.float(), weights.float() + + def sparse_matrix(self) -> torch.Tensor: + indices = torch.stack([self.row_idx, self.col_idx], dim=0) + return torch.sparse_coo_tensor( + indices=indices, + values=self.edge_weights, + size=(self.num_columns, self.num_columns), + device=self.edge_weights.device, + ).coalesce() + + def forward(self, states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Apply sparse inter-column message passing. + + Args: + states: `(B, C, D)` column states. + Returns: + inter_messages: `(B, C, D)` sparse messages. + scalar_activity: `(B, C)` pooled activity used for STDP. + """ + scalar_activity = states.mean(dim=-1) # (B, C) + sparse_w = self.sparse_matrix() + propagated = torch.sparse.mm(sparse_w, scalar_activity.transpose(0, 1)).transpose(0, 1) + inter_messages = propagated.unsqueeze(-1) * self.inter_gain.view(1, 1, -1) + return inter_messages, scalar_activity + + @torch.no_grad() + def stdp_update(self, pre_activity: torch.Tensor, post_activity: torch.Tensor) -> None: + """Hebbian/STDP-like update on sparse edge values.""" + if pre_activity.numel() == 0: + return + pre = pre_activity.mean(dim=0) # (C,) + post = post_activity.mean(dim=0) # (C,) + corr = torch.outer(post, pre) + edge_corr = corr[self.row_idx, self.col_idx] + anti = self.config.anti_hebbian * self.edge_weights.abs() + delta = self.config.stdp_lr * (edge_corr - anti) + new_w = self.edge_weights + delta + self.edge_weights.copy_(new_w.clamp(self.config.min_weight, self.config.max_weight)) + + def spatial_wiring_cost(self) -> torch.Tensor: + return (self.edge_weights.abs() * self.edge_distances).mean() + + def active_parameters(self) -> int: + return int((self.edge_weights.abs() > 1e-8).sum().item()) diff --git a/3D-SEMMN/semmn/dataset.py b/3D-SEMMN/semmn/dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..29c1fc64d3770dcab7f7f2a8e8b54cd2579265d0 --- /dev/null +++ b/3D-SEMMN/semmn/dataset.py @@ -0,0 +1,135 @@ +"""Multimodal AV-MNIST dataset helpers.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +from torch.utils.data import DataLoader, Dataset +from torchvision import datasets, transforms + +from semmn.config import SEMMNConfig + + +@dataclass +class Batch: + vision: torch.Tensor + audio: torch.Tensor + label: torch.Tensor + + +class SyntheticAVMNIST(Dataset): + """Pair MNIST vision with synthetic spectrogram-like audio.""" + + def __init__(self, root: str, train: bool = True) -> None: + tfm = transforms.ToTensor() + self.mnist = datasets.MNIST(root=root, train=train, download=True, transform=tfm) + + def __len__(self) -> int: + return len(self.mnist) + + def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + vision, label = self.mnist[idx] + audio = self._label_to_spectrogram(int(label), idx) + return vision, audio, torch.tensor(label, dtype=torch.long) + + def _label_to_spectrogram(self, label: int, seed: int) -> torch.Tensor: + g = torch.Generator().manual_seed(seed) + t = torch.linspace(0, 1, steps=224) + base_freq = 110.0 + 22.0 * float(label) + sine = ( + torch.sin(2.0 * torch.pi * base_freq * t) + + 0.3 * torch.sin(2.0 * torch.pi * 2.0 * base_freq * t) + + 0.2 * torch.randn(224, generator=g) + ) + spec = torch.stft( + sine, + n_fft=64, + hop_length=16, + win_length=64, + return_complex=True, + ).abs() + spec = spec.unsqueeze(0) # (1, F, T) + spec = F.interpolate(spec.unsqueeze(0), size=(112, 112), mode="bilinear", align_corners=False) + spec = spec.squeeze(0) + spec = spec / (spec.max().clamp_min(1e-6)) + return spec + + +class HuggingFaceAVMNIST(Dataset): + """Attempt to read BLOSSOM-framework/AV-MNIST with robust key fallback.""" + + def __init__(self, split: str = "train") -> None: + from datasets import load_dataset # lazy import + + self.ds = load_dataset("BLOSSOM-framework/AV-MNIST", split=split) + + def __len__(self) -> int: + return len(self.ds) + + def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + row = self.ds[idx] + vision = _to_tensor_image(row.get("image") or row.get("vision") or row.get("mnist")) + audio = _to_tensor_audio(row.get("audio") or row.get("spectrogram") or row.get("sound")) + label = int(row.get("label") or row.get("digit") or row.get("class")) + return vision, audio, torch.tensor(label, dtype=torch.long) + + +def _to_tensor_image(value) -> torch.Tensor: + if isinstance(value, torch.Tensor): + x = value.float() + else: + x = torch.tensor(value, dtype=torch.float32) + if x.ndim == 2: + x = x.unsqueeze(0) + x = F.interpolate(x.unsqueeze(0), size=(28, 28), mode="bilinear", align_corners=False).squeeze(0) + return x.clamp(0.0, 1.0) + + +def _to_tensor_audio(value) -> torch.Tensor: + if isinstance(value, dict) and "array" in value: + arr = torch.tensor(value["array"], dtype=torch.float32) + spec = torch.stft(arr, n_fft=128, hop_length=32, return_complex=True).abs() + x = spec.unsqueeze(0) + elif isinstance(value, torch.Tensor): + x = value.float() + else: + x = torch.tensor(value, dtype=torch.float32) + while x.ndim < 3: + x = x.unsqueeze(0) + if x.shape[0] != 1: + x = x[:1] + x = F.interpolate(x.unsqueeze(0), size=(112, 112), mode="bilinear", align_corners=False).squeeze(0) + x = x / (x.max().clamp_min(1e-6)) + return x + + +def build_dataloaders(config: SEMMNConfig) -> tuple[DataLoader, DataLoader]: + root = "./data" + if config.training.dataset_source == "hf": + try: + train_ds = HuggingFaceAVMNIST(split="train") + val_ds = HuggingFaceAVMNIST(split="test") + except Exception: + train_ds = SyntheticAVMNIST(root=root, train=True) + val_ds = SyntheticAVMNIST(root=root, train=False) + else: + train_ds = SyntheticAVMNIST(root=root, train=True) + val_ds = SyntheticAVMNIST(root=root, train=False) + + train_loader = DataLoader( + train_ds, + batch_size=config.training.batch_size, + shuffle=True, + num_workers=config.training.num_workers, + pin_memory=True, + ) + val_loader = DataLoader( + val_ds, + batch_size=config.training.batch_size, + shuffle=False, + num_workers=max(1, config.training.num_workers // 2), + pin_memory=True, + ) + return train_loader, val_loader diff --git a/3D-SEMMN/semmn/encoders.py b/3D-SEMMN/semmn/encoders.py new file mode 100644 index 0000000000000000000000000000000000000000..55b17d707cfdd928dc34147a3942f1acb2c9e907 --- /dev/null +++ b/3D-SEMMN/semmn/encoders.py @@ -0,0 +1,66 @@ +"""Vision/audio encoders and lobe injection.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + + +class VisionEncoder(nn.Module): + def __init__(self, out_dim: int = 256) -> None: + super().__init__() + self.net = nn.Sequential( + nn.Conv2d(1, 32, kernel_size=3, padding=1), + nn.ReLU(inplace=True), + nn.MaxPool2d(2), + nn.Conv2d(32, 64, kernel_size=3, padding=1), + nn.ReLU(inplace=True), + nn.MaxPool2d(2), + nn.Flatten(), + nn.Linear(64 * 7 * 7, out_dim), + nn.ReLU(inplace=True), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.net(x) + + +class AudioEncoder1D(nn.Module): + """1D Conv encoder over spectrogram time sequences.""" + + def __init__(self, out_dim: int = 256, freq_bins: int = 112) -> None: + super().__init__() + self.net = nn.Sequential( + nn.Conv1d(freq_bins, 128, kernel_size=5, padding=2), + nn.ReLU(inplace=True), + nn.Conv1d(128, 128, kernel_size=5, padding=2), + nn.ReLU(inplace=True), + nn.AdaptiveAvgPool1d(1), + nn.Flatten(), + nn.Linear(128, out_dim), + nn.ReLU(inplace=True), + ) + + def forward(self, spec: torch.Tensor) -> torch.Tensor: + # Input expected `(B, 1, F, T)`; use `(B, F, T)` for Conv1d. + x = spec.squeeze(1) + return self.net(x) + + +class LobeInjector(nn.Module): + """Project modality features into selected columns.""" + + def __init__(self, in_dim: int, num_columns: int, state_dim: int, target_columns: torch.Tensor) -> None: + super().__init__() + self.num_columns = num_columns + self.state_dim = state_dim + self.register_buffer("target_columns", target_columns.long()) + out_dim = max(1, self.target_columns.numel()) * state_dim + self.project = nn.Linear(in_dim, out_dim) + + def forward(self, features: torch.Tensor) -> torch.Tensor: + b = features.shape[0] + injected = features.new_zeros(b, self.num_columns, self.state_dim) + projected = self.project(features).view(b, -1, self.state_dim) + injected[:, self.target_columns, :] = projected[:, : self.target_columns.numel(), :] + return injected diff --git a/3D-SEMMN/semmn/grid.py b/3D-SEMMN/semmn/grid.py new file mode 100644 index 0000000000000000000000000000000000000000..a9eeac65d93b96850f2deabcd3bc4aae995f71da --- /dev/null +++ b/3D-SEMMN/semmn/grid.py @@ -0,0 +1,100 @@ +"""3D neuron grid utilities for SEMMN.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from semmn.config import SEMMNConfig + + +@dataclass +class GridPartition: + coordinates: torch.Tensor + column_ids: torch.Tensor + column_centers: torch.Tensor + column_sizes: torch.Tensor + column_to_neuron_indices: list[torch.Tensor] + vision_mask: torch.Tensor + audio_mask: torch.Tensor + hub_mask: torch.Tensor + vision_columns: torch.Tensor + audio_columns: torch.Tensor + hub_columns: torch.Tensor + + +def build_spatial_grid(config: SEMMNConfig, device: torch.device | None = None) -> GridPartition: + """Create fixed 3D Euclidean coordinates and column/lobe assignments.""" + gx, gy, gz = config.grid.dims + nx, ny = config.grid.n_columns_xy + device = device or torch.device("cpu") + + x = torch.arange(gx, dtype=torch.float32) + y = torch.arange(gy, dtype=torch.float32) + z = torch.arange(gz, dtype=torch.float32) + mesh = torch.meshgrid(x, y, z, indexing="ij") + coordinates = torch.stack(mesh, dim=-1).reshape(-1, 3).to(device) + + x_bins = torch.clamp((coordinates[:, 0] / max(gx, 1) * nx).long(), max=nx - 1) + y_bins = torch.clamp((coordinates[:, 1] / max(gy, 1) * ny).long(), max=ny - 1) + column_ids = x_bins * ny + y_bins + n_columns = nx * ny + + column_centers = torch.zeros(n_columns, 3, device=device) + column_sizes = torch.zeros(n_columns, device=device) + column_to_neuron_indices: list[torch.Tensor] = [] + for cid in range(n_columns): + idx = torch.where(column_ids == cid)[0] + column_to_neuron_indices.append(idx) + if idx.numel() > 0: + column_centers[cid] = coordinates[idx].mean(dim=0) + column_sizes[cid] = float(idx.numel()) + else: + xi = cid // ny + yi = cid % ny + column_centers[cid] = torch.tensor( + [(xi + 0.5) * gx / nx, (yi + 0.5) * gy / ny, 0.5 * gz], device=device + ) + + y_norm = coordinates[:, 1] / max(float(gy), 1.0) + z_norm = coordinates[:, 2] / max(float(gz), 1.0) + vision_mask = y_norm < config.grid.vision_y_ratio + audio_mask = y_norm >= config.grid.audio_y_ratio + hub_mask = ( + (y_norm >= config.grid.hub_y_range[0]) + & (y_norm <= config.grid.hub_y_range[1]) + & (z_norm >= config.grid.hub_z_range[0]) + & (z_norm <= config.grid.hub_z_range[1]) + ) + + vision_columns = _columns_from_mask(column_ids, vision_mask, n_columns) + audio_columns = _columns_from_mask(column_ids, audio_mask, n_columns) + hub_columns = _columns_from_mask(column_ids, hub_mask, n_columns) + + if vision_columns.numel() == 0: + vision_columns = torch.arange(0, n_columns // 3, device=device) + if audio_columns.numel() == 0: + audio_columns = torch.arange((2 * n_columns) // 3, n_columns, device=device) + if hub_columns.numel() == 0: + start = n_columns // 3 + hub_columns = torch.arange(start, start + max(1, n_columns // 5), device=device) + + return GridPartition( + coordinates=coordinates, + column_ids=column_ids, + column_centers=column_centers, + column_sizes=column_sizes, + column_to_neuron_indices=column_to_neuron_indices, + vision_mask=vision_mask, + audio_mask=audio_mask, + hub_mask=hub_mask, + vision_columns=vision_columns, + audio_columns=audio_columns, + hub_columns=hub_columns, + ) + + +def _columns_from_mask(column_ids: torch.Tensor, mask: torch.Tensor, n_columns: int) -> torch.Tensor: + hits = torch.bincount(column_ids[mask], minlength=n_columns) + return torch.where(hits > 0)[0] diff --git a/3D-SEMMN/semmn/hub.py b/3D-SEMMN/semmn/hub.py new file mode 100644 index 0000000000000000000000000000000000000000..8097d5d89a74a0258f8dd9eadfdadd3a249a0725 --- /dev/null +++ b/3D-SEMMN/semmn/hub.py @@ -0,0 +1,92 @@ +"""Shared manifold hub with cross-modal attention.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + + +class SharedManifoldHub(nn.Module): + def __init__( + self, + state_dim: int, + latent_dim: int, + feature_dim: int, + attention_heads: int = 4, + ) -> None: + super().__init__() + self.state_dim = state_dim + self.modality_proj = nn.Linear(feature_dim, state_dim) + self.attn = nn.MultiheadAttention( + embed_dim=state_dim, num_heads=attention_heads, batch_first=True + ) + self.hub_to_latent = nn.Linear(state_dim, latent_dim) + + self.reliability = nn.Sequential( + nn.Linear(feature_dim * 2, 64), + nn.ReLU(inplace=True), + nn.Linear(64, 2), + ) + + self.audio_imagination_decoder = nn.Sequential( + nn.Linear(latent_dim, feature_dim), + nn.ReLU(inplace=True), + nn.Linear(feature_dim, feature_dim), + ) + self.vision_imagination_decoder = nn.Sequential( + nn.Linear(latent_dim, feature_dim), + nn.ReLU(inplace=True), + nn.Linear(feature_dim, feature_dim), + ) + + self.cycle_vision = nn.Sequential( + nn.Linear(feature_dim, feature_dim), + nn.ReLU(inplace=True), + nn.Linear(feature_dim, feature_dim), + ) + self.cycle_audio = nn.Sequential( + nn.Linear(feature_dim, feature_dim), + nn.ReLU(inplace=True), + nn.Linear(feature_dim, feature_dim), + ) + + def forward( + self, + states: torch.Tensor, + hub_columns: torch.Tensor, + vision_features: torch.Tensor, + audio_features: torch.Tensor, + ) -> dict[str, torch.Tensor]: + hub_states = states[:, hub_columns, :] + v_token = self.modality_proj(vision_features).unsqueeze(1) + a_token = self.modality_proj(audio_features).unsqueeze(1) + tokens = torch.cat([hub_states, v_token, a_token], dim=1) + + attn_out, attn_map = self.attn(tokens, tokens, tokens, need_weights=True) + pooled = attn_out.mean(dim=1) + hub_latent = self.hub_to_latent(pooled) + + # Reliability-weighted fusion of modality latents. + rel_logits = self.reliability(torch.cat([vision_features, audio_features], dim=-1)) + rel_weights = torch.softmax(rel_logits, dim=-1) + fused_latent = ( + rel_weights[:, 0:1] * self.hub_to_latent(self.modality_proj(vision_features)) + + rel_weights[:, 1:2] * self.hub_to_latent(self.modality_proj(audio_features)) + + hub_latent + ) / 2.0 + + audio_from_vision = self.audio_imagination_decoder(fused_latent) + vision_from_audio = self.vision_imagination_decoder(fused_latent) + cycle_vision = self.cycle_vision(vision_from_audio) + cycle_audio = self.cycle_audio(audio_from_vision) + + return { + "hub_latent": hub_latent, + "fused_latent": fused_latent, + "reliability_weights": rel_weights, + "attention_map": attn_map, + "audio_from_vision": audio_from_vision, + "vision_from_audio": vision_from_audio, + "cycle_vision": cycle_vision, + "cycle_audio": cycle_audio, + } diff --git a/3D-SEMMN/semmn/losses.py b/3D-SEMMN/semmn/losses.py new file mode 100644 index 0000000000000000000000000000000000000000..9a72895f28b5fe89e28d8bee7ce31346f3fe1a59 --- /dev/null +++ b/3D-SEMMN/semmn/losses.py @@ -0,0 +1,60 @@ +"""Losses for multimodal manifold alignment.""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from semmn.config import SEMMNConfig +from semmn.connectivity import SparseInterColumnConnectivity + + +def nt_xent_loss(x: torch.Tensor, y: torch.Tensor, temperature: float = 0.1) -> torch.Tensor: + x = F.normalize(x, dim=-1) + y = F.normalize(y, dim=-1) + logits = x @ y.transpose(0, 1) / max(temperature, 1e-6) + labels = torch.arange(x.shape[0], device=x.device) + return 0.5 * (F.cross_entropy(logits, labels) + F.cross_entropy(logits.transpose(0, 1), labels)) + + +def compute_losses( + outputs: dict[str, torch.Tensor], + labels: torch.Tensor, + config: SEMMNConfig, + connectivity: SparseInterColumnConnectivity, +) -> dict[str, torch.Tensor]: + cls_loss = F.cross_entropy(outputs["logits"], labels) + contrastive = nt_xent_loss( + outputs["vision_embed"], outputs["audio_embed"], config.loss.temperature + ) + recon = F.mse_loss(outputs["audio_from_vision"], outputs["audio_features"].detach()) + F.mse_loss( + outputs["vision_from_audio"], outputs["vision_features"].detach() + ) + imagination = F.mse_loss( + outputs["cycle_vision"], outputs["vision_features"].detach() + ) + F.mse_loss(outputs["cycle_audio"], outputs["audio_features"].detach()) + + manifold = outputs["manifold_loss"] + spatial = connectivity.spatial_wiring_cost() + if not config.loss.use_manifold_loss: + manifold = manifold * 0.0 + if not config.loss.use_spatial_penalty: + spatial = spatial * 0.0 + + total = ( + cls_loss + + config.loss.contrastive_weight * contrastive + + config.loss.reconstruction_weight * recon + + config.loss.imagination_weight * imagination + + config.loss.manifold_weight * manifold + + config.loss.spatial_penalty_weight * spatial + ) + return { + "total": total, + "classification": cls_loss, + "contrastive": contrastive, + "reconstruction": recon, + "imagination": imagination, + "manifold": manifold, + "spatial": spatial, + } diff --git a/3D-SEMMN/semmn/manifold.py b/3D-SEMMN/semmn/manifold.py new file mode 100644 index 0000000000000000000000000000000000000000..b97748a0317de9f9400d16974d162095efcacdcd --- /dev/null +++ b/3D-SEMMN/semmn/manifold.py @@ -0,0 +1,71 @@ +"""Manifold regularization and geometric twist layer.""" + +from __future__ import annotations + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class GeometricTwistLayer(nn.Module): + """Learnable orthogonal transform via skew-symmetric matrix exponential. + + We use `exp(A - A^T)` instead of Cayley solve for MPS stability while + preserving an orthogonal rotation on the latent manifold. + """ + + def __init__(self, dim: int) -> None: + super().__init__() + self.a = nn.Parameter(torch.zeros(dim, dim)) + + def rotation_matrix(self) -> torch.Tensor: + skew = self.a - self.a.transpose(0, 1) + return torch.matrix_exp(skew) + + def forward(self, z: torch.Tensor) -> torch.Tensor: + return z @ self.rotation_matrix().transpose(0, 1) + + +class PopulationManifoldAE(nn.Module): + """Low-dimensional manifold constraint on column-pooled activity. + + In line with Gallego et al. (2017), this encourages dynamics to occupy + low-dimensional trajectories rather than unconstrained high-D activity. + """ + + def __init__(self, num_columns: int, state_dim: int, latent_dim: int = 8) -> None: + super().__init__() + self.num_columns = num_columns + self.state_dim = state_dim + self.latent_dim = latent_dim + + in_dim = num_columns + self.encoder = nn.Sequential( + nn.Linear(in_dim, 64), + nn.ReLU(inplace=True), + nn.Linear(64, 32), + nn.ReLU(inplace=True), + nn.Linear(32, latent_dim), + ) + self.decoder = nn.Sequential( + nn.Linear(latent_dim, 32), + nn.ReLU(inplace=True), + nn.Linear(32, 64), + nn.ReLU(inplace=True), + nn.Linear(64, in_dim), + ) + self.twist = GeometricTwistLayer(latent_dim) + + def forward(self, states: torch.Tensor) -> dict[str, torch.Tensor]: + pooled = states.mean(dim=-1) + latent = self.encoder(pooled) + twisted = self.twist(latent) + recon = self.decoder(twisted) + recon_loss = F.mse_loss(recon, pooled) + return { + "pooled_activity": pooled, + "latent": latent, + "twisted_latent": twisted, + "manifold_recon": recon, + "manifold_loss": recon_loss, + } diff --git a/3D-SEMMN/semmn/model.py b/3D-SEMMN/semmn/model.py new file mode 100644 index 0000000000000000000000000000000000000000..1db87871c37e3e1670bdd89c8435adf760d968ce --- /dev/null +++ b/3D-SEMMN/semmn/model.py @@ -0,0 +1,124 @@ +"""Top-level SEMMN model.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + +from semmn.columns import ColumnarDynamics +from semmn.config import SEMMNConfig +from semmn.connectivity import SparseInterColumnConnectivity +from semmn.encoders import AudioEncoder1D, LobeInjector, VisionEncoder +from semmn.grid import GridPartition, build_spatial_grid +from semmn.hub import SharedManifoldHub +from semmn.manifold import PopulationManifoldAE + + +class SEMMNModel(nn.Module): + """3D Spatially-Embedded Multimodal Manifold Network (prototype).""" + + def __init__(self, config: SEMMNConfig, device: torch.device | None = None) -> None: + super().__init__() + self.config = config + self.grid: GridPartition = build_spatial_grid(config, device=torch.device("cpu")) + self.num_columns = config.n_columns + self.state_dim = config.model.state_dim + + self.vision_encoder = VisionEncoder(config.model.vision_feature_dim) + self.audio_encoder = AudioEncoder1D(config.model.audio_feature_dim) + + self.vision_injector = LobeInjector( + in_dim=config.model.vision_feature_dim, + num_columns=self.num_columns, + state_dim=self.state_dim, + target_columns=self.grid.vision_columns, + ) + self.audio_injector = LobeInjector( + in_dim=config.model.audio_feature_dim, + num_columns=self.num_columns, + state_dim=self.state_dim, + target_columns=self.grid.audio_columns, + ) + + self.connectivity = SparseInterColumnConnectivity( + column_centers=self.grid.column_centers, + state_dim=self.state_dim, + config=config.connectivity, + ) + self.columnar = ColumnarDynamics( + num_columns=self.num_columns, state_dim=self.state_dim, rank=8 + ) + self.manifold = PopulationManifoldAE( + num_columns=self.num_columns, + state_dim=self.state_dim, + latent_dim=config.model.latent_dim, + ) + self.hub = SharedManifoldHub( + state_dim=self.state_dim, + latent_dim=config.model.latent_dim, + feature_dim=config.model.vision_feature_dim, + attention_heads=config.model.attention_heads, + ) + + self.vision_proj = nn.Linear(config.model.vision_feature_dim, config.model.latent_dim) + self.audio_proj = nn.Linear(config.model.audio_feature_dim, config.model.latent_dim) + self.classifier = nn.Linear(config.model.latent_dim, 10) + + def forward(self, vision: torch.Tensor, audio: torch.Tensor) -> dict[str, torch.Tensor]: + b = vision.shape[0] + device = vision.device + + vision_features = self.vision_encoder(vision) + audio_features = self.audio_encoder(audio) + + state = torch.zeros(b, self.num_columns, self.state_dim, device=device) + state = state + self.vision_injector(vision_features) + self.audio_injector(audio_features) + + pre_activity = None + post_activity = None + for _ in range(self.config.model.recurrent_steps): + inter_messages, pre_activity = self.connectivity(state) + state = self.columnar(state, inter_messages) + post_activity = state.mean(dim=-1) + + manifold_out = self.manifold(state) + hub_out = self.hub( + states=state, + hub_columns=self.grid.hub_columns.to(device), + vision_features=vision_features, + audio_features=audio_features, + ) + + vision_embed = self.vision_proj(vision_features) + audio_embed = self.audio_proj(audio_features) + fused = (hub_out["fused_latent"] + manifold_out["twisted_latent"]) / 2.0 + logits = self.classifier(fused) + + return { + "logits": logits, + "vision_embed": vision_embed, + "audio_embed": audio_embed, + "vision_features": vision_features, + "audio_features": audio_features, + "manifold_loss": manifold_out["manifold_loss"], + "latent": manifold_out["twisted_latent"], + "hub_latent": hub_out["hub_latent"], + "audio_from_vision": hub_out["audio_from_vision"], + "vision_from_audio": hub_out["vision_from_audio"], + "cycle_vision": hub_out["cycle_vision"], + "cycle_audio": hub_out["cycle_audio"], + "pre_activity": pre_activity, + "post_activity": post_activity, + "column_states": state, + "reliability_weights": hub_out["reliability_weights"], + } + + @torch.no_grad() + def apply_stdp(self, pre_activity: torch.Tensor, post_activity: torch.Tensor) -> None: + self.connectivity.stdp_update(pre_activity, post_activity) + + def active_parameters(self) -> int: + return self.connectivity.active_parameters() + + def energy_proxy(self) -> torch.Tensor: + return self.connectivity.spatial_wiring_cost() diff --git a/3D-SEMMN/train.py b/3D-SEMMN/train.py new file mode 100644 index 0000000000000000000000000000000000000000..46e70cdfdfb1052151d82338e55a7a88479e352b --- /dev/null +++ b/3D-SEMMN/train.py @@ -0,0 +1,154 @@ +"""Training entrypoint for 3D-SEMMN.""" + +from __future__ import annotations + +import argparse +import json +import random +from pathlib import Path + +import numpy as np +import torch +from torch.amp import GradScaler, autocast + +from semmn.config import SEMMNConfig +from semmn.dataset import build_dataloaders +from semmn.losses import compute_losses +from semmn.model import SEMMNModel + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Train 3D-SEMMN prototype") + parser.add_argument("--config", type=str, default="config/light.yaml") + parser.add_argument("--device", type=str, default="cuda") + parser.add_argument("--output-dir", type=str, default=None) + return parser.parse_args() + + +def set_seed(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +@torch.no_grad() +def evaluate( + model: SEMMNModel, + loader, + cfg: SEMMNConfig, + device: torch.device, + max_batches: int, +) -> dict[str, float]: + model.eval() + total = 0 + correct = 0 + retrieval_hits = 0 + retrieval_total = 0 + for bi, (vision, audio, label) in enumerate(loader): + if bi >= max_batches: + break + vision = vision.to(device, non_blocking=True) + audio = audio.to(device, non_blocking=True) + label = label.to(device, non_blocking=True) + out = model(vision, audio) + pred = out["logits"].argmax(dim=1) + correct += (pred == label).sum().item() + total += label.numel() + + v = torch.nn.functional.normalize(out["vision_embed"], dim=-1) + a = torch.nn.functional.normalize(out["audio_embed"], dim=-1) + sim = v @ a.t() + retrieved = sim.argmax(dim=1) + retrieval_hits += (label == label[retrieved]).sum().item() + retrieval_total += label.numel() + + return { + "accuracy": correct / max(total, 1), + "cross_modal_retrieval": retrieval_hits / max(retrieval_total, 1), + } + + +def main() -> None: + args = parse_args() + cfg = SEMMNConfig.from_yaml(args.config) + set_seed(cfg.seed) + + device = torch.device(args.device if torch.cuda.is_available() else "cpu") + out_dir = Path(args.output_dir or cfg.paths.output_dir) + out_dir.mkdir(parents=True, exist_ok=True) + + train_loader, val_loader = build_dataloaders(cfg) + model = SEMMNModel(cfg).to(device) + optimizer = torch.optim.AdamW( + model.parameters(), lr=cfg.training.lr, weight_decay=cfg.training.weight_decay + ) + scaler = GradScaler(enabled=(cfg.training.amp and device.type == "cuda")) + + history: list[dict[str, float]] = [] + best_retrieval = -1.0 + global_step = 0 + + for epoch in range(cfg.training.epochs): + model.train() + for step, (vision, audio, label) in enumerate(train_loader): + vision = vision.to(device, non_blocking=True) + audio = audio.to(device, non_blocking=True) + label = label.to(device, non_blocking=True) + + optimizer.zero_grad(set_to_none=True) + with autocast(device_type=device.type, enabled=scaler.is_enabled()): + outputs = model(vision, audio) + losses = compute_losses(outputs, label, cfg, model.connectivity) + + scaler.scale(losses["total"]).backward() + scaler.step(optimizer) + scaler.update() + + global_step += 1 + if global_step % cfg.connectivity.stdp_interval == 0: + model.apply_stdp(outputs["pre_activity"], outputs["post_activity"]) + + if step % cfg.training.log_interval == 0: + row = { + "epoch": float(epoch), + "step": float(step), + "loss_total": float(losses["total"].detach().cpu()), + "loss_cls": float(losses["classification"].detach().cpu()), + "loss_contrastive": float(losses["contrastive"].detach().cpu()), + "loss_manifold": float(losses["manifold"].detach().cpu()), + "active_params": float(model.active_parameters()), + "energy_proxy": float(model.energy_proxy().detach().cpu()), + } + history.append(row) + print(json.dumps(row)) + + eval_metrics = evaluate(model, val_loader, cfg, device, cfg.training.val_batches) + history.append( + { + "epoch": float(epoch), + "val_accuracy": float(eval_metrics["accuracy"]), + "val_cross_modal_retrieval": float(eval_metrics["cross_modal_retrieval"]), + } + ) + print(json.dumps({"epoch": epoch, **eval_metrics})) + + if eval_metrics["cross_modal_retrieval"] > best_retrieval: + best_retrieval = eval_metrics["cross_modal_retrieval"] + torch.save( + { + "model": model.state_dict(), + "config": cfg.__dict__, + "epoch": epoch, + "metrics": eval_metrics, + }, + out_dir / "best.pt", + ) + + with (out_dir / "history.json").open("w", encoding="utf-8") as f: + json.dump(history, f, indent=2) + torch.save({"model": model.state_dict()}, out_dir / "last.pt") + + +if __name__ == "__main__": + main() diff --git a/3D-SEMMN/visualize.py b/3D-SEMMN/visualize.py new file mode 100644 index 0000000000000000000000000000000000000000..f4d05d120f0425534cb95a89f3c88f689a69471f --- /dev/null +++ b/3D-SEMMN/visualize.py @@ -0,0 +1,503 @@ +"""Visualization utilities for 3D-SEMMN outputs.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import torch +import plotly.graph_objects as go +from sklearn.decomposition import PCA + +from semmn.config import SEMMNConfig +from semmn.dataset import build_dataloaders +from semmn.model import SEMMNModel + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Visualize 3D-SEMMN outputs") + parser.add_argument("--config", type=str, default="config/light.yaml") + parser.add_argument("--checkpoint", type=str, default="outputs/best.pt") + parser.add_argument("--history", type=str, default="outputs/history.json") + parser.add_argument("--output-dir", type=str, default="outputs/figures") + parser.add_argument("--device", type=str, default="cuda") + return parser.parse_args() + + +def plot_neuron_clusters(model: SEMMNModel, output_dir: Path) -> None: + coords = model.grid.coordinates.cpu().numpy() + columns = model.grid.column_ids.cpu().numpy() + + # Downsample for readability. + max_points = 30000 + if coords.shape[0] > max_points: + idx = np.random.choice(coords.shape[0], max_points, replace=False) + coords = coords[idx] + columns = columns[idx] + + fig = plt.figure(figsize=(9, 7)) + ax = fig.add_subplot(111, projection="3d") + ax.scatter(coords[:, 0], coords[:, 1], coords[:, 2], c=columns, s=1, cmap="tab20") + ax.set_title("3D Neuron Grid (colored by column)") + ax.set_xlabel("X") + ax.set_ylabel("Y") + ax.set_zlabel("Z") + fig.tight_layout() + fig.savefig(output_dir / "neuron_clusters_3d.png", dpi=180) + plt.close(fig) + + +def _sample_grid_dataframe(model: SEMMNModel, max_points: int = 40000) -> pd.DataFrame: + coords = model.grid.coordinates.cpu().numpy() + vision = model.grid.vision_mask.cpu().numpy() + audio = model.grid.audio_mask.cpu().numpy() + hub = model.grid.hub_mask.cpu().numpy() + columns = model.grid.column_ids.cpu().numpy() + + if coords.shape[0] > max_points: + idx = np.random.default_rng(42).choice(coords.shape[0], max_points, replace=False) + coords = coords[idx] + vision = vision[idx] + audio = audio[idx] + hub = hub[idx] + columns = columns[idx] + + module = np.full(coords.shape[0], "other", dtype=object) + module[audio] = "audio" + module[vision] = "vision" + module[hub] = "hub" # hub has highest priority if overlapping masks + return pd.DataFrame( + { + "x": coords[:, 0], + "y": coords[:, 1], + "z": coords[:, 2], + "column": columns, + "module": module, + } + ) + + +def plot_module_focus_images(model: SEMMNModel, output_dir: Path) -> None: + df = _sample_grid_dataframe(model, max_points=30000) + colors = {"vision": "#1f77b4", "audio": "#ff7f0e", "hub": "#2ca02c", "other": "#9467bd"} + modules = ["vision", "audio", "hub", "other"] + + # Baseline all-modules image. + fig = plt.figure(figsize=(10, 8)) + ax = fig.add_subplot(111, projection="3d") + for m in modules: + sub = df[df["module"] == m] + ax.scatter(sub["x"], sub["y"], sub["z"], s=1, alpha=0.65, c=colors[m], label=m) + ax.set_title("Neuron Grid by Module") + ax.set_xlabel("X") + ax.set_ylabel("Y") + ax.set_zlabel("Z") + ax.legend(loc="upper right") + fig.tight_layout() + fig.savefig(output_dir / "neuron_modules_all.png", dpi=180) + plt.close(fig) + + # Focus images with non-focused modules dimmed. + for focus in modules: + fig = plt.figure(figsize=(10, 8)) + ax = fig.add_subplot(111, projection="3d") + for m in modules: + sub = df[df["module"] == m] + if m == focus: + ax.scatter(sub["x"], sub["y"], sub["z"], s=2, alpha=0.95, c=colors[m], label=f"{m} (focus)") + else: + ax.scatter(sub["x"], sub["y"], sub["z"], s=1, alpha=0.05, c="#A0A0A0", label=f"{m} (dim)") + ax.set_title(f"Neuron Grid Focus: {focus}") + ax.set_xlabel("X") + ax.set_ylabel("Y") + ax.set_zlabel("Z") + ax.legend(loc="upper right") + fig.tight_layout() + fig.savefig(output_dir / f"neuron_module_{focus}.png", dpi=180) + plt.close(fig) + + +def plot_interactive_module_grid(model: SEMMNModel, output_dir: Path) -> None: + """Interactive HTML with module selection and dimming.""" + df = _sample_grid_dataframe(model, max_points=35000) + colors = {"vision": "#1f77b4", "audio": "#ff7f0e", "hub": "#2ca02c", "other": "#9467bd"} + modules = ["vision", "audio", "hub", "other"] + + fig = go.Figure() + for m in modules: + sub = df[df["module"] == m] + fig.add_trace( + go.Scatter3d( + x=sub["x"], + y=sub["y"], + z=sub["z"], + mode="markers", + name=m, + marker={"size": 2, "color": colors[m], "opacity": 0.75}, + customdata=np.stack([sub["column"].to_numpy(), sub["module"].to_numpy()], axis=1), + hovertemplate="x=%{x}
y=%{y}
z=%{z}
column=%{customdata[0]}
module=%{customdata[1]}", + ) + ) + + def focus_opacity(focus: str | None) -> list[float]: + if focus is None: + return [0.75, 0.75, 0.75, 0.75] + return [0.95 if m == focus else 0.06 for m in modules] + + buttons = [ + { + "label": "All modules", + "method": "restyle", + "args": [{"marker.opacity": focus_opacity(None)}], + } + ] + for m in modules: + buttons.append( + { + "label": f"Focus: {m}", + "method": "restyle", + "args": [{"marker.opacity": focus_opacity(m)}], + } + ) + + fig.update_layout( + title="3D-SEMMN Neuron Grid (Interactive Module Focus)", + scene={"xaxis_title": "X", "yaxis_title": "Y", "zaxis_title": "Z"}, + updatemenus=[ + { + "type": "dropdown", + "x": 0.01, + "y": 1.08, + "xanchor": "left", + "yanchor": "top", + "buttons": buttons, + } + ], + legend={"orientation": "h", "x": 0.01, "y": 0.98}, + margin={"l": 0, "r": 0, "t": 70, "b": 0}, + ) + fig.write_html(str(output_dir / "neuron_grid_interactive.html"), include_plotlyjs="cdn") + + +def _column_module_labels(model: SEMMNModel) -> np.ndarray: + n_columns = model.num_columns + labels = np.full(n_columns, "other", dtype=object) + labels[model.grid.audio_columns.cpu().numpy()] = "audio" + labels[model.grid.vision_columns.cpu().numpy()] = "vision" + labels[model.grid.hub_columns.cpu().numpy()] = "hub" + return labels + + +def _simulate_column_route( + start_activity: torch.Tensor, + row_idx: torch.Tensor, + col_idx: torch.Tensor, + weights: torch.Tensor, + steps: int, + top_edges_per_step: int = 80, +) -> tuple[pd.DataFrame, pd.DataFrame]: + """Simulate signal routing over learned sparse graph. + + This approximates route-of-understanding by propagating activity over + inter-column edges and keeping strongest contributors each recurrent step. + """ + act = start_activity.detach().float().clamp_min(0.0) + act = act / (act.sum().clamp_min(1e-6)) + + edge_rows: list[dict[str, float]] = [] + node_rows: list[dict[str, float]] = [] + + for step in range(steps): + node_strength = act.detach().cpu().numpy() + for node_id, val in enumerate(node_strength): + if val > 1e-4: + node_rows.append({"step": float(step), "column": float(node_id), "activity": float(val)}) + + edge_strength = weights * act[col_idx] + if edge_strength.numel() == 0: + break + k = min(top_edges_per_step, edge_strength.numel()) + top_vals, top_idx = torch.topk(edge_strength, k=k, largest=True) + for i in range(k): + s = int(col_idx[top_idx[i]].item()) + t = int(row_idx[top_idx[i]].item()) + w = float(top_vals[i].item()) + if w > 0: + edge_rows.append( + {"step": float(step), "source": float(s), "target": float(t), "strength": w} + ) + + next_act = torch.zeros_like(act) + next_act.index_add_(0, row_idx, edge_strength) + # Small residual keeps originating evidence traceable in later steps. + act = 0.8 * next_act + 0.2 * act + act = act / (act.sum().clamp_min(1e-6)) + + edge_df = pd.DataFrame(edge_rows) if edge_rows else pd.DataFrame(columns=["step", "source", "target", "strength"]) + node_df = pd.DataFrame(node_rows) if node_rows else pd.DataFrame(columns=["step", "column", "activity"]) + return edge_df, node_df + + +@torch.no_grad() +def plot_understanding_routes_2d( + model: SEMMNModel, cfg: SEMMNConfig, device: torch.device, output_dir: Path +) -> None: + """2D route maps of how visual/audio evidence traverses column graph.""" + _, val_loader = build_dataloaders(cfg) + batch = next(iter(val_loader)) + vision, audio, label = batch + vision = vision[:1].to(device) + audio = audio[:1].to(device) + digit = int(label[0].item()) + + vision_feat = model.vision_encoder(vision) + audio_feat = model.audio_encoder(audio) + vision_init = model.vision_injector(vision_feat)[0].norm(dim=-1) + audio_init = model.audio_injector(audio_feat)[0].norm(dim=-1) + + row_idx = model.connectivity.row_idx + col_idx = model.connectivity.col_idx + weights = model.connectivity.edge_weights.clamp_min(0.0) + centers = model.grid.column_centers[:, :2].detach().cpu().numpy() + modules = _column_module_labels(model) + module_colors = {"vision": "#1f77b4", "audio": "#ff7f0e", "hub": "#2ca02c", "other": "#A8A8A8"} + + vision_edges, vision_nodes = _simulate_column_route( + vision_init, row_idx, col_idx, weights, steps=cfg.model.recurrent_steps + ) + audio_edges, audio_nodes = _simulate_column_route( + audio_init, row_idx, col_idx, weights, steps=cfg.model.recurrent_steps + ) + + vision_edges["modality"] = "vision" + audio_edges["modality"] = "audio" + edge_frames = [df for df in (vision_edges, audio_edges) if not df.empty] + route_edges = pd.concat(edge_frames, ignore_index=True) if edge_frames else pd.DataFrame() + vision_nodes["modality"] = "vision" + audio_nodes["modality"] = "audio" + node_frames = [df for df in (vision_nodes, audio_nodes) if not df.empty] + route_nodes = pd.concat(node_frames, ignore_index=True) if node_frames else pd.DataFrame() + route_edges.to_csv(output_dir / "route_edges.csv", index=False) + route_nodes.to_csv(output_dir / "route_nodes.csv", index=False) + + def draw_static(modality: str, edge_df: pd.DataFrame, node_df: pd.DataFrame) -> None: + fig, ax = plt.subplots(figsize=(10, 8)) + ax.scatter(centers[:, 0], centers[:, 1], s=14, c="#E0E0E0", alpha=0.35, label="all columns") + for module_name, color in module_colors.items(): + idx = np.where(modules == module_name)[0] + if idx.size > 0: + ax.scatter( + centers[idx, 0], centers[idx, 1], s=18, c=color, alpha=0.30, label=f"{module_name} columns" + ) + + if not edge_df.empty: + max_strength = max(float(edge_df["strength"].max()), 1e-6) + for _, row in edge_df.iterrows(): + s = int(row["source"]) + t = int(row["target"]) + strength = float(row["strength"]) + alpha = min(0.95, 0.08 + 0.9 * (strength / max_strength)) + width = 0.7 + 2.4 * (strength / max_strength) + ax.plot( + [centers[s, 0], centers[t, 0]], + [centers[s, 1], centers[t, 1]], + color=module_colors[modality], + alpha=alpha, + linewidth=width, + ) + + if not node_df.empty: + final_step = int(node_df["step"].max()) + final_nodes = node_df[node_df["step"] == final_step] + top_nodes = final_nodes.nlargest(12, "activity") + idx = top_nodes["column"].astype(int).to_numpy() + ax.scatter( + centers[idx, 0], + centers[idx, 1], + s=120, + c=module_colors[modality], + edgecolors="black", + linewidths=0.8, + label=f"{modality} dominant columns", + ) + + hub_idx = np.where(modules == "hub")[0] + if hub_idx.size > 0: + ax.scatter( + centers[hub_idx, 0], + centers[hub_idx, 1], + s=40, + facecolors="none", + edgecolors="#2ca02c", + linewidths=1.0, + label="hub columns", + ) + + ax.set_title(f"2D Understanding Route ({modality}) - sample digit {digit}") + ax.set_xlabel("Column X") + ax.set_ylabel("Column Y") + ax.legend(loc="upper right", fontsize=8) + ax.grid(alpha=0.15) + fig.tight_layout() + fig.savefig(output_dir / f"route_map_{modality}_2d.png", dpi=200) + plt.close(fig) + + draw_static("vision", vision_edges, vision_nodes) + draw_static("audio", audio_edges, audio_nodes) + + # Interactive route map with modality focus selector. + fig = go.Figure() + fig.add_trace( + go.Scatter( + x=centers[:, 0], + y=centers[:, 1], + mode="markers", + marker={"size": 7, "color": "#D0D0D0", "opacity": 0.35}, + name="all_columns", + hovertemplate="column=%{text}", + text=[str(i) for i in range(centers.shape[0])], + ) + ) + + def edge_trace(df: pd.DataFrame, color: str, name: str) -> go.Scatter: + x_vals: list[float | None] = [] + y_vals: list[float | None] = [] + for _, row in df.iterrows(): + s = int(row["source"]) + t = int(row["target"]) + x_vals.extend([float(centers[s, 0]), float(centers[t, 0]), None]) + y_vals.extend([float(centers[s, 1]), float(centers[t, 1]), None]) + return go.Scatter( + x=x_vals, + y=y_vals, + mode="lines", + line={"width": 2, "color": color}, + opacity=0.75, + name=name, + hoverinfo="skip", + ) + + fig.add_trace(edge_trace(vision_edges, "#1f77b4", "vision_route")) + fig.add_trace(edge_trace(audio_edges, "#ff7f0e", "audio_route")) + fig.update_layout( + title=f"2D Route Simulation Across Column Graph (digit {digit})", + xaxis_title="Column X", + yaxis_title="Column Y", + updatemenus=[ + { + "type": "dropdown", + "x": 0.01, + "y": 1.12, + "buttons": [ + {"label": "Show both", "method": "update", "args": [{"visible": [True, True, True]}]}, + {"label": "Vision route", "method": "update", "args": [{"visible": [True, True, False]}]}, + {"label": "Audio route", "method": "update", "args": [{"visible": [True, False, True]}]}, + ], + } + ], + legend={"orientation": "h"}, + ) + fig.write_html(str(output_dir / "route_map_interactive.html"), include_plotlyjs="cdn") + + +@torch.no_grad() +def plot_manifold_trajectories( + model: SEMMNModel, cfg: SEMMNConfig, device: torch.device, output_dir: Path +) -> None: + _, val_loader = build_dataloaders(cfg) + zs = [] + ys = [] + for bi, (vision, audio, label) in enumerate(val_loader): + if bi >= 20: + break + out = model(vision.to(device), audio.to(device)) + zs.append(out["latent"].cpu()) + ys.append(label) + latent = torch.cat(zs, dim=0).numpy() + labels = torch.cat(ys, dim=0).numpy() + proj = PCA(n_components=2).fit_transform(latent) + + fig, ax = plt.subplots(figsize=(7, 6)) + scatter = ax.scatter(proj[:, 0], proj[:, 1], c=labels, cmap="tab10", s=8) + ax.set_title("Manifold Trajectories (PCA of latent-8)") + ax.set_xlabel("PC1") + ax.set_ylabel("PC2") + fig.colorbar(scatter, ax=ax, label="Digit") + fig.tight_layout() + fig.savefig(output_dir / "manifold_pca.png", dpi=180) + plt.close(fig) + + +def plot_metrics(history_path: Path, output_dir: Path) -> None: + with history_path.open("r", encoding="utf-8") as f: + rows = json.load(f) + + train_rows = [r for r in rows if "loss_total" in r] + val_rows = [r for r in rows if "val_accuracy" in r] + if not train_rows: + return + + fig, axes = plt.subplots(1, 3, figsize=(15, 4)) + axes[0].plot([r["loss_total"] for r in train_rows], label="total") + axes[0].plot([r["loss_cls"] for r in train_rows], label="cls") + axes[0].set_title("Training Loss") + axes[0].legend() + + axes[1].plot([r["energy_proxy"] for r in train_rows], label="energy") + axes[1].plot([r["active_params"] for r in train_rows], label="active_params") + axes[1].set_title("Energy/Active Params") + axes[1].legend() + + if val_rows: + axes[2].plot([r["val_accuracy"] for r in val_rows], label="val_acc") + axes[2].plot( + [r["val_cross_modal_retrieval"] for r in val_rows], + label="cross_modal", + ) + axes[2].set_title("Validation Metrics") + axes[2].legend() + fig.tight_layout() + fig.savefig(output_dir / "metrics.png", dpi=180) + plt.close(fig) + + +def main() -> None: + args = parse_args() + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + + cfg = SEMMNConfig.from_yaml(args.config) + if args.device == "mps": + use_device = "mps" if torch.backends.mps.is_available() else "cpu" + elif args.device == "cuda": + use_device = "cuda" if torch.cuda.is_available() else "cpu" + else: + use_device = args.device + device = torch.device(use_device) + model = SEMMNModel(cfg).to(device) + + ckpt = torch.load(args.checkpoint, map_location=device) + model_state = model.state_dict() + incoming = ckpt["model"] + filtered = { + k: v + for k, v in incoming.items() + if k in model_state and model_state[k].shape == v.shape + } + model.load_state_dict(filtered, strict=False) + model.eval() + + plot_neuron_clusters(model, output_dir) + plot_module_focus_images(model, output_dir) + plot_interactive_module_grid(model, output_dir) + plot_understanding_routes_2d(model, cfg, device, output_dir) + plot_manifold_trajectories(model, cfg, device, output_dir) + plot_metrics(Path(args.history), output_dir) + + +if __name__ == "__main__": + main() diff --git a/3D_SEMMN_final.pdf b/3D_SEMMN_final.pdf new file mode 100644 index 0000000000000000000000000000000000000000..6f64abc3f1cfac0380291c8dc50b584bc7390b15 --- /dev/null +++ b/3D_SEMMN_final.pdf @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b3a00719a292a9411573cc83d78fb9a22ef60ba546e556deeb659efc97902688 +size 159578