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