Inspect Hidden-State Propagation
A neural cellular automaton can solve a maze while most of its computation remains invisible.
The visible output might be one distance channel.
The internal state may contain eleven hidden channels (5–15 in the maze layout) evolving underneath it.
This chapter asks:
what is moving through those hidden channels while the answer is being computed?
Capture the entire state trajectory
Instead of saving only the final output — rollout_frozen plus per-step recording, same clamped execution, now with a trace:
@torch.no_grad()
def trace_rollout(model, state, steps, frozen_inputs):
trace = [state.detach().cpu()]
for _ in range(steps):
state = model(state)
state[:, :3] = frozen_inputs[:, :3]
trace.append(state.detach().cpu())
return torch.stack(trace)
The resulting tensor has a conceptual shape like:
time × batch × channel × height × width
Now the recurrent computation is data we can analyze.
Plot one hidden channel through time
import matplotlib.pyplot as plt
def show_channel(trace, channel, times):
fig, axes = plt.subplots(1, len(times), figsize=(3 * len(times), 3))
for ax, t in zip(axes, times):
ax.imshow(trace[t, 0, channel], cmap="coolwarm")
ax.set_title(f"t={t}")
ax.axis("off")
plt.tight_layout()
Try:
show_channel(trace, channel=7, times=[0, 4, 8, 16, 32, 64])
Some channels may look like noise.
Others may show spatial waves, boundary responses or persistent local markers.
Do not assign semantic names too quickly.
Measure where a channel becomes active
A simple activity map:
def channel_activity(trace, channel, threshold=0.1):
values = trace[:, 0, channel].abs()
return (values > threshold).float().mean(dim=0)
This answers:
which cells used this channel frequently?
Compare activity with:
walls
frontiers
goal distance
branch points
final path
Spatial alignment can suggest hypotheses.
It does not prove function.
Track information arrival time
For each cell, record when a hidden channel first crosses a threshold.
def first_activation_time(trace, channel, threshold=0.1):
active = trace[:, 0, channel].abs() > threshold
times = torch.full(active.shape[1:], -1, dtype=torch.long)
for t in range(active.shape[0]):
new = active[t] & (times < 0)
times[new] = t
return times
Plot that map.
If activation time grows with distance from the goal or start, the channel may participate in a propagating signal.
That is much more informative than one final heatmap.
One locality bound frames every such map: with a 3×3 neighborhood, influence spreads at most one cell per step, so arrival times earlier than distance-from-source (in cells) rule out local propagation — and would reveal positional leakage instead. The bound is architecture, not behavior:
after t steps,
influence radius ≤ neighborhood radius × t
Compare hidden channels with BFS quantities
We have exact classical reference signals available:
distance from goal
distance from start
reachable mask
BFS frontier arrival time
shortest-path membership
For each hidden channel, compute simple correlations.
def correlation(a, b, mask=None):
if mask is not None:
a = a[mask]
b = b[mask]
a = a.float().flatten()
b = b.float().flatten()
a = a - a.mean()
b = b - b.mean()
return (a * b).mean() / (a.std(unbiased=False) * b.std(unbiased=False) + 1e-8)
(Verified: self-correlation reads exactly 1.0; constant inputs read 0.0.)
A channel strongly correlated with BFS distance is interesting.
But correlation still does not mean the model explicitly represents “distance” in that channel.
signal propagation
≠ symbolic reasoning
hidden-channel structure
≠ interpretable algorithm
Probe hidden state with a simple decoder
Freeze the NCA.
Collect hidden states from many mazes.
Then train a small linear probe to predict a known quantity such as BFS distance.
Conceptually:
probe = torch.nn.Conv2d(hidden_channels, 1, kernel_size=1)
Only train the probe.
If a linear decoder can recover distance, then distance-related information is accessible in the hidden representation.
Again, be precise:
linearly decodable
is not the same as:
used causally by the NCA
Look at temporal phase changes
The hidden computation may not have one stationary meaning.
A channel can behave differently during:
early propagation
mid-rollout conflict resolution
late stabilization
So compute statistics by time window:
def temporal_energy(trace, channel):
x = trace[:, 0, channel]
return x.pow(2).mean(dim=(1, 2))
Plot energy versus step.
A channel that peaks early and disappears may be carrying transient frontier information.
A channel that remains active may encode persistent structure.
Visualize gradients too
Another question is:
which cells and channels can influence the final decision?
Keep one rollout differentiable and backpropagate from the output at a selected location.
state.requires_grad_(True)
final = rollout(model, state, steps=64)
score = final[0, 3, query_y, query_x]
score.backward()
influence = state.grad.abs().sum(dim=1)[0]
This gives a local sensitivity map of the initial state.
For recurrent systems, such maps should be interpreted cautiously: gradients can vanish, explode or reflect only local linear sensitivity around one trajectory.
Still, they provide another view.
Hidden states are not explanations by themselves
A colorful channel is easy to narrate.
That is dangerous.
A responsible workflow is — each rung strictly stronger evidence than the last:
flowchart LR
O[observe pattern] --> H[hypothesis + BFS comparison]
H --> P[probe: linear decoder]
P --> I[intervene on state]
I --> M[measure behavioral change]
| Rung | Method | Evidence strength | Establishes at most |
|---|---|---|---|
| observe | channel plots, arrival maps | suggestive | where activity concentrates |
| correlate | BFS-distance correlation | associative | statistical dependence |
| probe | frozen linear decoder | representational | decodability, not causality |
| intervene | zero/shuffle/freeze/swap | causal | necessity under that intervention |
The intervention step is crucial.
That is what we do next.
In the final NCA chapter we will zero, shuffle, freeze and perturb hidden channels to ask which internal signals are actually necessary for the learned computation.
Research
Earle, S., Yildiz, O., Togelius, J. & Hegde, C. — Pathfinding Neural Cellular Automata (2023). The reference signals this chapter correlates against: hand-coded BFS/DFS wavefronts as the classical procedures the hidden dynamics are tested against — arrival times, frontiers, and distances with known ground truth. https://arxiv.org/abs/2301.06820
Mordvintsev, A., Randazzo, E., Niklasson, E. & Levin, M. — Growing Neural Cellular Automata (Distill, 2020). The hidden-state framing owned here: twelve channels without predefined meaning, offered as chemical-signaling analogy — with this chapter’s ladder (observe → correlate → probe → intervene) as the discipline that keeps the analogy honest. https://doi.org/10.23915/distill.00023