Graph World Models: The Technicalities
The notation behind graph world models: the transition function, fixed versus dynamic edges, and where the training objective comes from.
As I said in the main article, a world model’s functioning is very simple:
where represents the state of the world at time , represents the action, and is the world model itself.
Now graph world models are no different. We have
where the only difference is that represents the state of the world as a graph:
where is the vertex set, the adjacency at time , and the node features. In practice you store the edge list , since takes up quadratic space.
That leaves , which the equation assumes and never explains. In the simplest formulation the action is a global vector that conditions the whole transition, so every node’s update sees the same action:
For the power grid below, that vector encodes something like open breaker 7. For the molecule example, raise the temperature by 10K. For the coding agent in the case study, it is a one-hot over which tool fired and which file it touched. The action does not have to be global. You can also inject it at a single node and let message passing carry the effect outward, which is closer to what the agentic case really does, but the global form is the one to keep in your head.
There are two paradigms of graph world models:
- Fixed edge GWMs
- Dynamic edge GWMs
Fixed edge GWMs
In this version the graph topology is fixed. The support of the adjacency, meaning which pairs of nodes are connected at all, never changes across the rollout. What does change is the node features , and the weights sitting on that fixed support.
That distinction is easy to blur, so let me state it plainly: “the adjacency stays the same” is not quite right, which is why still carries a . The wiring is fixed. The numbers on the wires are not.
A power grid is the clean example. Vertices are substations and generators. Edges are actual physical connections, copper in the ground that nobody is rewiring between timesteps, so the support is fixed. holds the normalised power flowing to neighbours through those edges, which changes every timestep. is the state of each node: temperature, load, phase, voltage.
Dynamic edge GWMs
Here the graph topology can change with each pass to the world model. This makes sense in autoregressive world models. Adjacency changes with each pass based on the changes in the node features.
Concretely:
where are the state representations of the nodes at some time .
Example could be molecule/atomic interactions. New interactions may arise depending on the state: temperature, magnetic field, electric field, pressure and so on. The world model could be used to simulate and study these interactions.
What is then?
We have covered what the input looks like in GWMs. Now for the mysterious black box. Alright, I am hyping this up to be very complicated, but it’s actually pretty simple. is any model that can take the graph as input and give the next graph state as output.
What really holds is the power of representing the graph as vectors or numbers that machines can understand. Improving the model performance implies getting better representations of the graph state. A weak world model is one whose representation collapses states that actually behave differently, so it predicts the same future for both. A strong one keeps them apart.
Training objective
The objective is that we need to maximize the probability where belong to the training data. Assume is the training data. Sample from .
which is the same thing as
Now what does this mean intuitively?
gives us the next graph state . We need to make sure that the predicted is the most likely graph state as given in the training data.
Let’s take the case of fixed edge GWMs and look at node features .
The assumption is that a node’s features at depend on the node features at and the edge connections at . Note this is the whole feature matrix , not just that node’s own row. Message passing is the entire reason we bothered with a graph, so node ‘s future has to be allowed to depend on its neighbours’ present:
That product hides a second assumption: given the current graph and the action, the nodes are treated as independent of each other. Each node gets its own prediction and they are never asked to agree. This is what every per-node prediction head does, and it is fine for training on one step at a time, but it means the model gives you marginals instead of one coherent joint next-graph. Sample from it repeatedly and the little inconsistencies between nodes are precisely the thing that compounds over a rollout. More on that in the limitations.
So the objective function becomes:
Let’s make this concrete. So far is an abstract distribution, and the objective is unfalsifiable in the sense that you cannot code it. Pick a distribution and the whole thing collapses into something familiar. Say the model’s prediction for each node is a Gaussian with identity covariance, and note that it is the prediction that is Gaussian, not the graph:
So the network emits exactly one thing per node, a predicted mean. Substituting the Gaussian density, every factor outside the exponent is a constant that does not depend on , so it drops straight out of the argmin:
leaving
That is mean squared error over node features, and it is the whole loss. All the likelihood machinery above was building up to the thing everyone reaches for by default anyway, which I find reassuring: squared error is the maximum likelihood objective under one specific assumption, and now you know which assumption you are making when you type it.
It also differentiates without drama, which matters more than it sounds:
The gradient is just the residual, with nothing in a denominator, nothing that has to be kept positive, and nothing that explodes when the model gets confident.
This was my attempt at actually understanding a simplistic model of training a GWM. The assumption of the covariance being an identity limits the training objective. It essentially implies that all node feature dimensions
in the same manner. This is not true for the power grid example itself: the temperature wiggles differently from the voltage or power. But the main point here was to get an intuition of what exactly we are optimizing in the model. Ideally, we would include the variance in the training objective itself.
Comments
Commenting uses a GitHub account, via GitHub Discussions.