From 49ef1b9291374ceaf026599608e7c429fdd46fc6 Mon Sep 17 00:00:00 2001 From: Rosie Lickorish Date: Tue, 29 Sep 2026 20:04:55 +0100 Subject: [PATCH 1/2] Optional --output_path arg added to evaluate subcommand Signed-off-by: Rosie Lickorish --- README.md | 1 + gridfm_graphkit/__main__.py | 6 + gridfm_graphkit/cli.py | 5 +- tests/test_predict_embeddings.py | 369 +++++++++++++++++++++++++++++++ 4 files changed, 380 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index 954f118e..f3d8104b 100644 --- a/README.md +++ b/README.md @@ -243,6 +243,7 @@ gridfm_graphkit evaluate --config path/to/eval.yaml --model_path path/to/model.p | `--profiler` | `str` | Enable Lightning profiler (`simple`, `advanced`, `pytorch`). | `None` | | `--compute_dc_ac_metrics` | `flag` | Compute ground-truth AC/DC power balance metrics on the test split. | `False` | | `--save_output` | `flag` | Save predictions under MLflow artifacts (`.../artifacts/test`). For the PowerFlow task this writes `_predictions.parquet` (bus-level) and `_branch_predictions.parquet` (branch-level flows, thermal loading, and angle violations). | `False` | +| `--output_path` | `str` | Optional custom output directory to save evaluation predictions (when used with `--save_output`). Defaults to MLflow artifacts directory (`artifacts/test`). | `None` | | `--mp_context` | `str` | DataLoader multiprocessing start method (`spawn`, `fork`, `forkserver`). Defaults to PyTorch's automatic choice. On Linux, `spawn` is recommended for safety (CUDA + fork is unsafe); other choices emit a warning. | `None` | ### Example with saved normalizer stats diff --git a/gridfm_graphkit/__main__.py b/gridfm_graphkit/__main__.py index aa50278b..b8ae1952 100644 --- a/gridfm_graphkit/__main__.py +++ b/gridfm_graphkit/__main__.py @@ -376,6 +376,12 @@ def main(): "--converted_config_path", **_converted_config_path_kwargs, ) + evaluate_parser.add_argument( + "--output_path", + type=str, + default=None, + help="Optional custom output directory to save evaluation predictions. Defaults to MLflow artifacts/test directory.", + ) evaluate_parser.add_argument("--stream-partitions", **_stream_partitions_kwargs) # ---- PREDICT SUBCOMMAND ---- diff --git a/gridfm_graphkit/cli.py b/gridfm_graphkit/cli.py index f8fd0625..4cb53be1 100644 --- a/gridfm_graphkit/cli.py +++ b/gridfm_graphkit/cli.py @@ -438,7 +438,10 @@ def main_cli(args): if args.command == "predict": output_dir = args.output_path else: - output_dir = os.path.join(artifacts_dir, "test") + output_dir = getattr(args, "output_path", None) or os.path.join( + artifacts_dir, + "test", + ) os.makedirs(output_dir, exist_ok=True) first_prediction = predictions[0] if any(isinstance(value, dict) for value in first_prediction.values()): diff --git a/tests/test_predict_embeddings.py b/tests/test_predict_embeddings.py index 12c170bf..786ec9b0 100644 --- a/tests/test_predict_embeddings.py +++ b/tests/test_predict_embeddings.py @@ -1,3 +1,4 @@ +import os import numpy as np import torch import yaml @@ -115,6 +116,56 @@ def test_predict_parser_accepts_get_embeddings_flag() -> None: assert parsed_args.get_embeddings is True +def test_predict_parser_output_path_defaults_to_data() -> None: + """When --output_path is omitted, it defaults to 'data'.""" + test_argv = [ + "gridfm_graphkit", + "predict", + "--config", + "examples/config/HGNS_PF_118Bus.yaml", + "--model_path", + "tests/models/dummy_model.pt", + ] + + with ( + mock.patch("sys.argv", test_argv), + mock.patch( + "gridfm_graphkit.__main__.main_cli", + ) as mocked_main_cli, + ): + main() + + parsed_args = mocked_main_cli.call_args.args[0] + assert parsed_args.command == "predict" + assert parsed_args.output_path == "data" + + +def test_predict_parser_accepts_output_path_flag() -> None: + """--output_path is a registered optional argument on the predict subcommand.""" + test_argv = [ + "gridfm_graphkit", + "predict", + "--config", + "examples/config/HGNS_PF_118Bus.yaml", + "--model_path", + "tests/models/dummy_model.pt", + "--output_path", + "/some/custom/path", + ] + + with ( + mock.patch("sys.argv", test_argv), + mock.patch( + "gridfm_graphkit.__main__.main_cli", + ) as mocked_main_cli, + ): + main() + + parsed_args = mocked_main_cli.call_args.args[0] + assert parsed_args.command == "predict" + assert parsed_args.output_path == "/some/custom/path" + + def test_main_cli_propagates_get_embeddings_to_task_args(tmp_path) -> None: config_path = tmp_path / "config.yaml" config_path.write_text( @@ -578,6 +629,324 @@ def fake_to_parquet(df, path, index=False): assert all(not index for _, _, index in saved.values()) +# --------------------------------------------------------------------------- +# evaluate --output_path: parser registration and output-dir routing +# --------------------------------------------------------------------------- + + +def test_evaluate_parser_accepts_output_path_flag() -> None: + """--output_path is a registered argument on the evaluate subcommand.""" + test_argv = [ + "gridfm_graphkit", + "evaluate", + "--config", + "examples/config/HGNS_PF_118Bus.yaml", + "--model_path", + "tests/models/dummy_model.pt", + "--save_output", + "--output_path", + "/some/custom/path", + ] + + with ( + mock.patch("sys.argv", test_argv), + mock.patch( + "gridfm_graphkit.__main__.main_cli", + ) as mocked_main_cli, + ): + main() + + parsed_args = mocked_main_cli.call_args.args[0] + assert parsed_args.command == "evaluate" + assert parsed_args.output_path == "/some/custom/path" + + +def test_evaluate_parser_output_path_defaults_to_none() -> None: + """When --output_path is omitted, it defaults to None.""" + test_argv = [ + "gridfm_graphkit", + "evaluate", + "--config", + "examples/config/HGNS_PF_118Bus.yaml", + "--model_path", + "tests/models/dummy_model.pt", + ] + + with ( + mock.patch("sys.argv", test_argv), + mock.patch( + "gridfm_graphkit.__main__.main_cli", + ) as mocked_main_cli, + ): + main() + + parsed_args = mocked_main_cli.call_args.args[0] + assert parsed_args.command == "evaluate" + assert parsed_args.output_path is None + + +def test_main_cli_evaluate_uses_custom_output_path(tmp_path) -> None: + """When --output_path is set, predictions are written to that directory.""" + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "seed": 0, + "task": {"task_name": "PowerFlow"}, + "data": { + "networks": ["case14"], + "workers": 0, + "baseMVA": 100, + }, + "training": { + "accelerator": "cpu", + "devices": 1, + "strategy": "auto", + "epochs": 1, + "batch_size": 1, + }, + "callbacks": {"tol": 0.0, "patience": 1}, + "version": 1, + "optimizer": { + "type": "AdamW", + "learning_rate": 1e-3, + "optimizer_params": {"betas": [0.9, 0.999]}, + "scheduler_type": "ReduceLROnPlateau", + "scheduler_params": {"mode": "min", "factor": 0.5, "patience": 1}, + }, + }, + ), + ) + + custom_dir = str(tmp_path / "custom_output") + args = SimpleNamespace( + tf32=False, + log_dir=str(tmp_path / "mlruns"), + exp_name="tests", + run_name="evaluate-output-path", + config=str(config_path), + data_path=str(tmp_path / "data"), + model_path="tests/models/dummy_model.pt", + command="evaluate", + output_path=custom_dir, + get_embeddings=False, + num_workers=0, + batch_size=None, + plugins=[], + dataset_wrapper=None, + dataset_wrapper_cache_dir=None, + mp_context=None, + normalizer_stats=None, + bfloat16=False, + compile=None, + profiler=None, + report_performance=False, + deterministic=False, + compute_dc_ac_metrics=False, + save_output=True, + ) + + logger = SimpleNamespace( + save_dir=str(tmp_path / "mlruns"), + experiment_id="0", + run_id="abc", + ) + written_paths = [] + + class DummyModel: + def __init__(self): + self.model = self + + def load_state_dict(self, _state_dict): + return None + + class DummyDataModule: + def __init__(self, *args, **kwargs): + self.data_normalizers = [object()] + + class DummyTrainer: + def __init__(self, *args, **kwargs): + pass + + def test(self, model=None, datamodule=None): + return [{}] + + def predict(self, model=None, datamodule=None): + return [ + {"scenario": np.array([0]), "vm_pu": np.array([1.0])}, + {"scenario": np.array([1]), "vm_pu": np.array([1.1])}, + ] + + def fake_to_parquet(df, path, index=False): + written_paths.append(str(path)) + + with ( + mock.patch("gridfm_graphkit.cli.MLFlowLogger", return_value=logger), + mock.patch( + "gridfm_graphkit.cli.L.seed_everything", + ), + mock.patch( + "gridfm_graphkit.cli.LitGridHeteroDataModule", + DummyDataModule, + ), + mock.patch( + "gridfm_graphkit.cli.get_task", + return_value=DummyModel(), + ), + mock.patch( + "gridfm_graphkit.cli.L.Trainer", + DummyTrainer, + ), + mock.patch( + "gridfm_graphkit.cli.torch.load", + return_value={}, + ), + mock.patch( + "pandas.DataFrame.to_parquet", + new=fake_to_parquet, + ), + ): + main_cli(args) + + assert len(written_paths) == 1 + assert os.path.dirname(written_paths[0]) == custom_dir + + +def test_main_cli_evaluate_falls_back_to_artifacts_test_dir(tmp_path) -> None: + """When --output_path is omitted, predictions go to /test.""" + config_path = tmp_path / "config.yaml" + config_path.write_text( + yaml.safe_dump( + { + "seed": 0, + "task": {"task_name": "PowerFlow"}, + "data": { + "networks": ["case14"], + "workers": 0, + "baseMVA": 100, + }, + "training": { + "accelerator": "cpu", + "devices": 1, + "strategy": "auto", + "epochs": 1, + "batch_size": 1, + }, + "callbacks": {"tol": 0.0, "patience": 1}, + "version": 1, + "optimizer": { + "type": "AdamW", + "learning_rate": 1e-3, + "optimizer_params": {"betas": [0.9, 0.999]}, + "scheduler_type": "ReduceLROnPlateau", + "scheduler_params": {"mode": "min", "factor": 0.5, "patience": 1}, + }, + }, + ), + ) + + experiment_id = "0" + run_id = "xyz" + args = SimpleNamespace( + tf32=False, + log_dir=str(tmp_path / "mlruns"), + exp_name="tests", + run_name="evaluate-output-path", + config=str(config_path), + data_path=str(tmp_path / "data"), + model_path="tests/models/dummy_model.pt", + command="evaluate", + output_path=None, + get_embeddings=False, + num_workers=0, + batch_size=None, + plugins=[], + dataset_wrapper=None, + dataset_wrapper_cache_dir=None, + mp_context=None, + normalizer_stats=None, + bfloat16=False, + compile=None, + profiler=None, + report_performance=False, + deterministic=False, + compute_dc_ac_metrics=False, + save_output=True, + ) + + logger = SimpleNamespace( + save_dir=str(tmp_path / "mlruns"), + experiment_id=experiment_id, + run_id=run_id, + ) + expected_dir = os.path.join( + str(tmp_path / "mlruns"), + experiment_id, + run_id, + "artifacts", + "test", + ) + written_paths = [] + + class DummyModel: + def __init__(self): + self.model = self + + def load_state_dict(self, _state_dict): + return None + + class DummyDataModule: + def __init__(self, *args, **kwargs): + self.data_normalizers = [object()] + + class DummyTrainer: + def __init__(self, *args, **kwargs): + pass + + def test(self, model=None, datamodule=None): + return [{}] + + def predict(self, model=None, datamodule=None): + return [ + {"scenario": np.array([0]), "vm_pu": np.array([1.0])}, + {"scenario": np.array([1]), "vm_pu": np.array([1.1])}, + ] + + def fake_to_parquet(df, path, index=False): + written_paths.append(str(path)) + + with ( + mock.patch("gridfm_graphkit.cli.MLFlowLogger", return_value=logger), + mock.patch( + "gridfm_graphkit.cli.L.seed_everything", + ), + mock.patch( + "gridfm_graphkit.cli.LitGridHeteroDataModule", + DummyDataModule, + ), + mock.patch( + "gridfm_graphkit.cli.get_task", + return_value=DummyModel(), + ), + mock.patch( + "gridfm_graphkit.cli.L.Trainer", + DummyTrainer, + ), + mock.patch( + "gridfm_graphkit.cli.torch.load", + return_value={}, + ), + mock.patch( + "pandas.DataFrame.to_parquet", + new=fake_to_parquet, + ), + ): + main_cli(args) + + assert len(written_paths) == 1 + assert os.path.dirname(written_paths[0]) == expected_dir + + class _Const(torch.nn.Module): """Returns a constant tensor of width ``out_dim``, one row per input row.""" From 54cddb25a553be91786be26f06f8ba115a6746cb Mon Sep 17 00:00:00 2001 From: Rosie Lickorish Date: Thu, 1 Oct 2026 10:00:20 +0100 Subject: [PATCH 2/2] Update quick_start readme with --ouput_path for evaluate Signed-off-by: Rosie Lickorish --- docs/quick_start/quick_start.md | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/quick_start/quick_start.md b/docs/quick_start/quick_start.md index 5e619525..790402e1 100644 --- a/docs/quick_start/quick_start.md +++ b/docs/quick_start/quick_start.md @@ -109,6 +109,7 @@ gridfm_graphkit evaluate --config path/to/eval.yaml --model_path path/to/model.p | `--profiler` | `str` | Enable Lightning profiler (`simple`, `advanced`, `pytorch`). | `None` | | `--compute_dc_ac_metrics` | `flag` | Compute ground-truth AC/DC power balance metrics on the test split. | `False` | | `--save_output` | `flag` | Save predictions under MLflow artifacts (`.../artifacts/test`). For the PowerFlow task this writes `_predictions.parquet` (bus-level) and `_branch_predictions.parquet` (branch-level flows, thermal loading, and angle violations). | `False` | +| `--output_path` | `str` | Optional custom output directory to save evaluation predictions (when used with `--save_output`). Defaults to MLflow artifacts directory (`artifacts/test`). | `None` | | `--mp_context` | `str` | DataLoader multiprocessing start method (`spawn`, `fork`, `forkserver`). Defaults to PyTorch's automatic choice. On Linux, `spawn` is recommended for safety (CUDA + fork is unsafe); other choices emit a warning. | `None` | ### Example with saved normalizer stats