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:

st+1=Fθ(st,at)s_{t+1} = F_\theta(s_t, a_t)

where sts_t represents the state of the world at time tt, ata_t represents the action, and FθF_\theta is the world model itself.

Now graph world models are no different. We have

Gt+1=Fθ(Gt,at)G_{t+1} = F_\theta(G_t, a_t)

where the only difference is that GtG_t represents the state of the world as a graph:

Gt=(V,At,Xt)G_t = (V, A_t, X_t)

where VV is the vertex set, AtA_t the adjacency at time tt, and XtX_t the node features. In practice you store the edge list EtE_t, since AtA_t takes up quadratic space.

That leaves ata_t, which the equation assumes and never explains. In the simplest formulation the action is a global vector at∈Rda_t \in \mathbb{R}^d that conditions the whole transition, so every node’s update sees the same action:

hi′=UPDATE(hi, AGGj∈N(i) hj, at)h_i' = \text{UPDATE}\Big(h_i,\ \text{AGG}_{j \in \mathcal{N}(i)}\, h_j,\ a_t\Big)

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:

  1. Fixed edge GWMs
  2. 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 XtX_t, 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 AtA_t still carries a tt. 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. AtA_t holds the normalised power flowing to neighbours through those edges, which changes every timestep. XtX_t 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:

A^ij=σ(hi⊤Q hj)\hat{A}_{ij} = \sigma\big(h_i^\top Q\, h_j\big)

where hi,hjh_i, h_j are the state representations of the nodes at some time tt.

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 FθF_\theta 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. FθF_\theta is any model that can take the graph as input and give the next graph state as output.

What FθF_\theta 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 P(Gt+1∣Gt,at)P(G_{t+1}|G_{t},a_{t}) where Gt+1,GtG_{t+1}, G_{t} belong to the training data. Assume D\mathcal{D} is the training data. Sample (Gt,at,Gt+1)(G_t, a_t, G_{t+1}) from D\mathcal{D}.

argmaxθ  E(Gt,at,Gt+1)∼D[log⁡Pθ(Gt+1∣Gt,at)]\underset{\theta}{\text{argmax}} \; \mathbb{E}_{(G_t, a_t, G_{t+1}) \sim \mathcal{D}} \left[ \log P_{\theta}(G_{t+1} | G_t, a_t) \right]

which is the same thing as

minθ  E(Gt,at,Gt+1)∼D[−log⁡Pθ(Gt+1∣Gt,at)]\underset{\theta}{\text{min}} \; \mathbb{E}_{(G_t, a_t, G_{t+1}) \sim \mathcal{D}} \left[ -\log P_{\theta}(G_{t+1} | G_t, a_t) \right]

Now what does this mean intuitively?

FθF_\theta gives us the next graph state G^t+1\hat{G}_{t+1}. We need to make sure that the G^t+1\hat{G}_{t+1} 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 XtX_t.

Xt=[Xt,1,Xt,2,…,Xt,n],Xt,i∈Rd    ∀iX_t = [X_{t,1}, X_{t,2}, \dots, X_{t,n}], \qquad X_{t,i} \in \mathbb{R}^d \;\; \forall i

The assumption is that a node’s features at t+1t+1 depend on the node features at tt and the edge connections at tt. Note this is the whole feature matrix XtX_t, not just that node’s own row. Message passing is the entire reason we bothered with a graph, so node ii‘s future has to be allowed to depend on its neighbours’ present:

P(Gt+1∣Gt,at)=∏i=1nP(Xt+1,i∣Xt,At,at)P(G_{t+1}|G_t,a_t) = \prod_{i=1}^{n} P(X_{t+1,i} \mid X_{t}, A_{t}, a_{t})

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 nn 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:

minθ  E(Gt,at,Gt+1)∼D[−∑i=1nlog⁡Pθ(Xt+1,i∣Xt,At,at)]\underset{\theta}{\text{min}} \; \mathbb{E}_{(G_t, a_t, G_{t+1}) \sim \mathcal{D}} \left[ -\sum_{i=1}^{n} \log P_{\theta}(X_{t+1,i} \mid X_{t}, A_{t}, a_t) \right]

Let’s make this concrete. So far PθP_\theta 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:

Xt+1,i∼N(μθ,i, I),μθ=Fθ(Xt,At,at)X_{t+1,i} \sim \mathcal{N}\big(\mu_{\theta,i},\ I\big), \qquad \mu_{\theta} = F_\theta(X_t, A_t, a_t)

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 θ\theta, so it drops straight out of the argmin:

−∑i=1nlog⁡Pθ(Xt+1,i∣⋅)=∑i=1n12∥Xt+1,i−μθ,i∥2  +  const-\sum_{i=1}^{n} \log P_{\theta}(X_{t+1,i} \mid \cdot) = \sum_{i=1}^{n} \tfrac{1}{2}\lVert X_{t+1,i} - \mu_{\theta,i} \rVert^{2} \;+\; \text{const}

leaving

minθ  E(Gt,at,Gt+1)∼D[∑i=1n∥Xt+1,i−μθ,i∥2]\underset{\theta}{\text{min}} \; \mathbb{E}_{(G_t, a_t, G_{t+1}) \sim \mathcal{D}} \left[ \sum_{i=1}^{n} \lVert X_{t+1,i} - \mu_{\theta,i} \rVert^{2} \right]

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:

∂∂μθ,i  12∥Xt+1,i−μθ,i∥2  =  −(Xt+1,i−μθ,i)\frac{\partial}{\partial \mu_{\theta,i}} \; \tfrac{1}{2}\lVert X_{t+1,i} - \mu_{\theta,i} \rVert^{2} \;=\; -\big(X_{t+1,i} - \mu_{\theta,i}\big)

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 A Wigglytuff standing still and wobbling in place. 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.