World Models: predict the next world state instead of the next frame
· 5 min read
Note: While writing this I discovered some overlapping work from Odyssey (Agora-1) and AlayaLab (MASS, PWM), but I'm posting this anyway as it may still be helpful for someone catching up on world models.
BEGIN
Recent advances in world models have inspired me to ramp up on the research field. After reading some papers, I had a question to explore:
- Why does world state get updated by what we see? Can't we learn a latent world state that updates over time without observation?
Why the answer might be important
Before we get to simulation theory (don't we need this to make earth-2 or earth-N?), I'd argue it's important for today's world models as well. Disentangling world state from rendering may allow for faster training of difficult concepts, such as how far a candle wick has burned down since it was lit. A latent world state may help capture things that are hard to parameterize, like the candle example or wind/air effects.
Another reason I believe in this direction is for the sake of scaling multiplayer experiences. If world state is independent, each player can query it for their camera pose and render downstream. This should be less computationally expensive and more stable for multiple simultaneous actions affecting world state.
Similar works
LatentSpatialMemory shows the benefit of neural state over explicit 3d mapping. This still updates world state from observation though:

NeuWorld uses a renderable implicit scene, but its transition is conditioned on future camera trajectory, current observation, and retrieved visual history:

PERSIST updates world state as a 3D voxel grid, then decodes camera pose and pixels, which is close to what we wanted. They still predict future state based on pixel observation though (blue square in the architecture image):

Programmable World Model separates state from rendering, but the state transitions are controlled by a world program. Great for verifiable world state, but my bet is that we still need implicit "world programs" to learn all the dynamics of the world:

Experiments
We start by generating a room full of boxes with two players in a game engine. For training we can generate all kinds of room layouts and let our game engine generate various action sequences. Action sequences are just keyboard presses (WASD for movement).
Experiment #1 is to simply learn state dynamics without rendering. We train W (world state) to directly output continuous coordinates, appearance, velocity for each object in the room. W is initialized from the game engine ground truth for a given room, we take the actions, and compare predicted world to ground truth.
Experiment #2 is to learn an implicit world state W.

An encoder reads the starting state from the game engine and compresses it to implicit vectors (8×48 or 16×64 slots). Transitions are conditioned only on the input actions, and a prediction head is asked to answer "what's at this point in space?" for any (x, y, z) where we have ground truth from the game engine to use as a loss.
This model struggled to learn any dynamics. Hypothesis on this failure is that the model simply took a shortcut. The world state included all the wall, floor, and object information as equal and the model learned that if it predicts no movement (or random movement), it's still mostly correct \_(ツ)_/
There are some potential optimizations here to test but I decided not to spend more time working with this readout loss. Let's move on and learn from other spatial embedding strategies.
Experiment #3 is to define world state as particles (like 3DGS). We still "cheat" on the initialization by taking positions and colors of objects and players from the game engine and assigning "moving" particles to them while leaving the rest as "static" particles.
We train a graph network to learn physics based on "Learning to Simulate Complex Physics with Graph Networks" (Sanchez-Gonzalez et al., DeepMind, 2020). The "message passing" in the diagram below just means that if I push an object from the right side, the right side particles will pass that message to its nearby neighbors, six times over. This lets all the particles of the box know they need to be moving.
Those physics update our particle world state, which later is rendered by our game engine by looking at where our objects and characters are in the updated particle world state.

Recap learnings from each experiment
#1: dynamics can be learned from actions alone, without rendering in the loop
#2: training an implicit world latent on point queries that mostly land on static objects teaches the model to ignore actions (shocker)
#3: a state with spatial structure (particles) works and can learn physics
Lets review our opening questions
Can we hold and transition state without observation? Yes! With the caveat that in these experiments we borrowed the starting state from the simulator.
Does the state have to be latent? Still unclear, but this post is getting too long. So far, the learned particle state was the only working representation, but more on this in the next blog posts.
Is this important for multiplayer? Probably. World state transitions ran 7x real-time on an H100, and we can let clients handle rendering. Need to test at scale.
What's next?
I'm planning to revisit world state options after upgrading our environment with real physics and photorealistic worlds.
I'm also launching CAUS - open source world models optimized for Apple Silicon.