Skip to content
Open
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
35 changes: 28 additions & 7 deletions gridfm_graphkit/datasets/graph_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion gridfm_graphkit/datasets/powergrid_hetero_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
Loading