Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions graph_weather/models/regional_forecast.py
Original file line number Diff line number Diff line change
Expand Up @@ -274,14 +274,17 @@ def forward(
nodes = torch.cat([features[i], regional_h3], dim=0)
nodes = self.node_encoder(nodes)
nodes, _ = self.encoder_gnn(nodes, enc_graph.edge_index, enc_edge_attr)
obs_features = nodes[:num_obs]
h3_features = nodes[num_obs:]

# Process: N rounds of H3 message passing
h3_features = self.processor(h3_features, lat_graph.edge_index, latent_edge_attr)

# Decode: H3 -> obs through reversed bipartite GNN
obs_placeholders = torch.zeros(num_obs, self.config.node_dim, device=features.device)
dec_nodes = torch.cat([obs_placeholders, h3_features], dim=0)
# Decode: H3 -> obs through reversed bipartite GNN.
# Seed the decoder's observation nodes with their own encoded features rather
# than zeros, so the predicted delta can specialize per observation instead of
# only seeing its cell's pooled feature (mirrors the StretchedForecaster fix).
dec_nodes = torch.cat([obs_features, h3_features], dim=0)
dec_nodes, _ = self.decoder_gnn(dec_nodes, dec_edge_index, dec_edge_attr)
obs_out = self.node_decoder(dec_nodes[:num_obs])
batch_outputs.append(obs_out)
Expand Down
22 changes: 22 additions & 0 deletions tests/test_regional_forecast.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,28 @@ def test_residual_connection():
assert torch.allclose(out, expected_residual, atol=1e-5)


def test_decoder_seeded_with_encoder_obs_features():
"""The decoder's observation nodes carry the encoder's per-obs features, not zeros.

Seeding them with zeros collapses the model toward persistence, since the predicted delta
then only sees cell-level context. This guards that regression.
"""
model = _small_config().build()
features = torch.randn(1, 5, 16)

captured = {}

def capture(_module, inputs, _output):
captured["dec_input"] = inputs[0].detach()

handle = model.decoder_gnn.register_forward_hook(capture)
model(features, _uk_latlons())
handle.remove()

obs_rows = captured["dec_input"][:5]
assert obs_rows.abs().sum() > 0


def _nudging_config():
"""Config with nudging enabled."""
return RegionalForecasterConfig(
Expand Down
Loading