diff --git a/gridfm_graphkit/datasets/graph_builder.py b/gridfm_graphkit/datasets/graph_builder.py index c8b489f5..12d85633 100644 --- a/gridfm_graphkit/datasets/graph_builder.py +++ b/gridfm_graphkit/datasets/graph_builder.py @@ -72,6 +72,32 @@ "Ytf_i", ] + COMMON_BRANCH_FEATURES +# Column indices whose *pre-mask, pre-normalisation* values are snapshotted into +# ``.static`` at build time and restored by RemovePFMask after inference (see +# gridfm_graphkit.datasets.masking.RemovePFMask). Single source of truth for both +# the build path (build_hetero_data) and the backfill path (backfill_static). +BUS_STATIC_COLS = [MIN_VM_H, MAX_VM_H, MIN_QG_H, MAX_QG_H, VN_KV] +BRANCH_STATIC_COLS = [ANG_MIN, ANG_MAX, RATE_A] + + +def backfill_static(data: HeteroData) -> None: + """Attach ``.static`` limit snapshots to graphs that lack them, in place. + + Graphs processed before ``.static`` was introduced carry the limit columns + only inside ``bus.x`` / branch ``edge_attr``. Reconstruct the snapshot from + those raw columns so the PF test/predict path (RemovePFMask) works on such + caches without reprocessing the dataset. + + Must run on the raw graph *before* normalisation and branch masking, so the + captured values match what ``build_hetero_data`` stores (raw, full edge set). + """ + bus = data["bus"] + if not hasattr(bus, "static"): + bus.static = bus.x[:, BUS_STATIC_COLS].clone() + branch = data["bus", "connects", "bus"] + if not hasattr(branch, "static"): + branch.static = branch.edge_attr[:, BRANCH_STATIC_COLS].clone() + def build_hetero_data( bus_df: pd.DataFrame, @@ -102,9 +128,7 @@ def build_hetero_data( # Bus nodes data["bus"].x = torch.tensor(bus_df[BUS_FEATURES].values, dtype=torch.float) - data["bus"].static = ( - data["bus"].x[:, [MIN_VM_H, MAX_VM_H, MIN_QG_H, MAX_QG_H, VN_KV]].clone() - ) + data["bus"].static = data["bus"].x[:, BUS_STATIC_COLS].clone() # Generator nodes gen_df = gen_df.reset_index(drop=True) @@ -144,10 +168,7 @@ def build_hetero_data( data["bus", "connects", "bus"].edge_index = edge_index data["bus", "connects", "bus"].edge_attr = edge_attr - data["bus", "connects", "bus"].static = edge_attr[ - :, - [ANG_MIN, ANG_MAX, RATE_A], - ].clone() + data["bus", "connects", "bus"].static = edge_attr[:, BRANCH_STATIC_COLS].clone() data["bus", "connects", "bus"].y = edge_y # Gen-Bus and Bus-Gen edges diff --git a/gridfm_graphkit/datasets/powergrid_hetero_dataset.py b/gridfm_graphkit/datasets/powergrid_hetero_dataset.py index a28b3ab6..ec554a98 100644 --- a/gridfm_graphkit/datasets/powergrid_hetero_dataset.py +++ b/gridfm_graphkit/datasets/powergrid_hetero_dataset.py @@ -8,7 +8,7 @@ from tqdm import tqdm from typing import Optional, Callable from torch_geometric.data import HeteroData -from gridfm_graphkit.datasets.graph_builder import build_hetero_data +from gridfm_graphkit.datasets.graph_builder import backfill_static, build_hetero_data class HeteroGridDatasetDisk(Dataset): @@ -346,5 +346,8 @@ def get(self, idx): raise IndexError(f"Data file {file_name} does not exist.") data_dict = torch.load(file_name, weights_only=True) data = HeteroData.from_dict(data_dict) + # Backfill .static for caches processed before it was added, before + # normalisation so the snapshot matches build_hetero_data (raw columns). + backfill_static(data) self.data_normalizer.transform(data=data) return data