kiruluta's picture
Upload folder using huggingface_hub
4fd79a1 verified
|
Raw History Blame Contribute Delete
6.04 kB

Spectral World Models

Andrew Kiruluta, UC Berkeley, June 2026

This repository is a compact PyTorch implementation of the manuscript architecture Spectral World Models: Multimodal Operator Learning in Hilbert Space for Imagination-Based Planning.

It implements a concrete Spectral World Model (SWM) rather than a broad design space:

  • fixed 2-level Haar wavelet image encoder and inverse spectral decoder,
  • DCT-style spectral text encoder over token-position embeddings,
  • low-rank multimodal fusion into a shared latent Hilbert-style state,
  • action-conditioned low-rank spectral transition operator,
  • predictive, reconstruction, latent consistency, text, and stability losses,
  • imagination rollout evaluation,
  • benchmark baselines and ablations.

The bundled benchmark is Mini Moving Shapes, a generated multimodal next-state prediction dataset. Each sequence contains a moving square or circle, a symbolic text description, and discrete action controls. The task is to predict the next image/text state and maintain stable multi-step rollouts.

Core idea: Fuse image and text modalities into a shared spectral latent and use low-rank spectral (or baseline) transition operators to predict next-frame images and text, enabling short rollouts for imagination-based evaluation.

Key components Models:

  • models.py β€” wavelet/DCT/FFT/conv encoders, decoders, and transition modules (low-rank spectral, Koopman, MLP, transformer) plus build_model.
  • Data: data.py β€” Mini Moving Shapes synthetic generator and TransitionDataset / SequenceDataset.
  • Training: train.py β€” training loop, validation/test evaluation, rollout evaluation, checkpointing, and metric export.
  • Metrics: metrics.py β€” MSE, PSNR, SSIM, Gaussian NLL, token accuracy, parameter counting.
  • Utilities: haar.py β€” dependency-free 2D Haar wavelet transform.

Repository layout

spectral-world-models/
β”œβ”€β”€ data/
β”‚   β”œβ”€β”€ generate_benchmark.py
β”‚   └── mini_moving_shapes.npz
β”œβ”€β”€ spectral_world_models/
β”‚   β”œβ”€β”€ data.py
β”‚   β”œβ”€β”€ haar.py
β”‚   β”œβ”€β”€ metrics.py
β”‚   β”œβ”€β”€ models.py
β”‚   └── train.py
β”œβ”€β”€ scripts/
β”‚   └── run_quick_benchmark.py
β”œβ”€β”€ results/
β”‚   β”œβ”€β”€ README.md
β”‚   β”œβ”€β”€ metrics.csv
β”‚   β”œβ”€β”€ metrics.json
β”‚   └── *_metrics.json / *_history.csv / *.pt
β”œβ”€β”€ tests/
β”‚   └── test_smoke.py
β”œβ”€β”€ pyproject.toml
└── requirements.txt

Installation

python -m venv .venv
source .venv/bin/activate
pip install -e .

Or, without editable installation:

pip install -r requirements.txt
export PYTHONPATH=$PWD

Generate the benchmark dataset

The dataset is already included at data/mini_moving_shapes.npz. To regenerate it:

PYTHONPATH=. python data/generate_benchmark.py

Default split sizes:

  • train: 256 sequences
  • validation: 64 sequences
  • test: 64 sequences
  • sequence length: 8
  • image size: 32x32 grayscale
  • actions: stay, up, down, left, right

Run the full quick benchmark

PYTHONPATH=. python scripts/run_quick_benchmark.py

This trains and evaluates:

  1. Dreamer-style latent baseline
  2. Transformer latent predictor baseline
  3. Koopman baseline
  4. Neural-operator / Fourier-feature baseline
  5. Spectral World Model (ours)
  6. SWM ablations:
    • no spectral state
    • no spectral transition
    • no text branch
    • no stability penalty

Train selected models

PYTHONPATH=. python -m spectral_world_models.train \
  --data data/mini_moving_shapes.npz \
  --out results \
  --epochs 3 \
  --batch-size 64 \
  --models swm dreamer transformer koopman neural_operator

Metrics

The repo reports:

  • training loss per epoch,
  • validation loss, PSNR, SSIM, NLL, text accuracy,
  • test loss, PSNR, SSIM, NLL, text accuracy,
  • rollout MSE over horizon 5,
  • average rollout step time.

The included results/ folder contains actual quick CPU benchmark metrics generated in the sandbox. These are meant as reproducibility smoke-test metrics, not final paper-grade multi-seed values.

Main quick benchmark results

Method Params Test PSNR ↑ Test SSIM ↑ Test NLL ↓ Rollout MSE@5 ↓ Time/step ms ↓
Dreamer-style latent model 497,682 14.571 0.072 0.0107 0.03490 0.0528
Transformer latent predictor 503,322 14.368 0.017 0.0687 0.03657 0.0290
Koopman baseline 465,306 14.619 0.090 -0.0027 0.03487 0.0170
Neural-operator baseline 403,794 14.534 0.058 0.0210 0.03511 0.0194
Spectral World Model (ours) 405,422 14.591 0.080 0.0050 0.03577 0.0343

Ablation quick benchmark results

Variant Params Test PSNR ↑ Test NLL ↓ Rollout MSE@5 ↓ Time/step ms ↓
Spectral World Model (ours) 405,422 14.591 0.0050 0.03577 0.0343
No spectral state 497,682 14.567 0.0117 0.03492 0.0184
No spectral transition 446,802 14.562 0.0131 0.03505 0.0318
No text branch 402,350 14.729 -0.0325 0.03508 0.0302
No stability penalty 405,422 14.433 0.0498 0.03609 0.0294

Important interpretation note

The quick benchmark is deliberately small so the repository can be run on CPU. It verifies that the architecture, baselines, ablations, dataset, training loop, validation loop, test loop, and rollout metrics are implemented end-to-end. It is not intended as a final NeurIPS-grade empirical study. For paper-quality claims, rerun with larger sequence sets, longer training, multiple seeds, and a harder image/control benchmark.

Test

PYTHONPATH=. python -m pytest -q

Expected result:

1 passed

This repo is based on the following citation. Please cite this reference when using this repo in your work:

https://www.researchgate.net/publication/404985395_Spectral_World_Models_Multimodal_Operator_Learning_in_Hilbert_Space_for_Imagination-Based_Planning