diff --git a/graph_weather/models/regional_forecast.py b/graph_weather/models/regional_forecast.py index 14a19634..a35d9f12 100644 --- a/graph_weather/models/regional_forecast.py +++ b/graph_weather/models/regional_forecast.py @@ -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) diff --git a/tests/test_regional_forecast.py b/tests/test_regional_forecast.py index 1d78d940..25fdccaf 100644 --- a/tests/test_regional_forecast.py +++ b/tests/test_regional_forecast.py @@ -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(