diff --git a/examples/property_prediction/README.md b/examples/property_prediction/README.md index 599d650..3232cf1 100644 --- a/examples/property_prediction/README.md +++ b/examples/property_prediction/README.md @@ -46,6 +46,20 @@ Fine-tuning the pre-trained model on OGBG-MolHIV: bash hiv_pre.sh ``` +Pretrained weights can also be loaded from a local Fairseq checkpoint or a +raw model state dictionary: + +```bash +fairseq-train ... \ + --pretrained-model-path /path/to/checkpoint.pt +``` + +`--pretrained-model-path` and `--pretrained-model-name` are mutually +exclusive. Use `--load-pretrained-model-output-layer` when the local +checkpoint's output layer should be retained. Only load checkpoints from +trusted sources. For distributed jobs, the path must be readable by every +worker. + #### OGBG-MolHIV Method | #params | test AUC (%)| --------------|---------|------------| diff --git a/graphormer/evaluate/evaluate.py b/graphormer/evaluate/evaluate.py index 32c079c..0ec29e7 100644 --- a/graphormer/evaluate/evaluate.py +++ b/graphormer/evaluate/evaluate.py @@ -28,7 +28,10 @@ def eval(args, use_pretrained, checkpoint_path=None, logger=None): # load checkpoint if use_pretrained: - model_state = load_pretrained_model(cfg.task.pretrained_model_name) + model_state = load_pretrained_model( + cfg.task.pretrained_model_name, + cfg.task.pretrained_model_path, + ) else: model_state = torch.load(checkpoint_path)["model"] model.load_state_dict( @@ -83,17 +86,22 @@ def eval(args, use_pretrained, checkpoint_path=None, logger=None): y_true = torch.Tensor(y_true) # evaluate pretrained models - if use_pretrained: - if cfg.task.pretrained_model_name == "pcqm4mv1_graphormer_base": - evaluator = ogb.lsc.PCQM4MEvaluator() - input_dict = {'y_pred': y_pred, 'y_true': y_true} - result_dict = evaluator.eval(input_dict) - logger.info(f'PCQM4Mv1Evaluator: {result_dict}') - elif cfg.task.pretrained_model_name == "pcqm4mv2_graphormer_base": - evaluator = ogb.lsc.PCQM4Mv2Evaluator() - input_dict = {'y_pred': y_pred, 'y_true': y_true} - result_dict = evaluator.eval(input_dict) - logger.info(f'PCQM4Mv2Evaluator: {result_dict}') + if ( + use_pretrained + and cfg.task.pretrained_model_name == "pcqm4mv1_graphormer_base" + ): + evaluator = ogb.lsc.PCQM4MEvaluator() + input_dict = {'y_pred': y_pred, 'y_true': y_true} + result_dict = evaluator.eval(input_dict) + logger.info(f'PCQM4Mv1Evaluator: {result_dict}') + elif ( + use_pretrained + and cfg.task.pretrained_model_name == "pcqm4mv2_graphormer_base" + ): + evaluator = ogb.lsc.PCQM4Mv2Evaluator() + input_dict = {'y_pred': y_pred, 'y_true': y_true} + result_dict = evaluator.eval(input_dict) + logger.info(f'PCQM4Mv2Evaluator: {result_dict}') else: if args.metric == "auc": auc = roc_auc_score(y_true, y_pred) @@ -116,7 +124,7 @@ def main(): ) args = options.parse_args_and_arch(parser, modify_parser=None) logger = logging.getLogger(__name__) - if args.pretrained_model_name != "none": + if args.pretrained_model_name != "none" or args.pretrained_model_path: eval(args, True, logger=logger) elif hasattr(args, "save_dir"): for checkpoint_fname in os.listdir(args.save_dir): diff --git a/graphormer/models/graphormer.py b/graphormer/models/graphormer.py index f952c69..b92eb75 100644 --- a/graphormer/models/graphormer.py +++ b/graphormer/models/graphormer.py @@ -39,9 +39,16 @@ def __init__(self, args, encoder): if getattr(args, "apply_graphormer_init", False): self.apply(init_graphormer_params) self.encoder_embed_dim = args.encoder_embed_dim - if args.pretrained_model_name != "none": - self.load_state_dict(load_pretrained_model(args.pretrained_model_name)) - if not args.load_pretrained_model_output_layer: + pretrained_model_name = getattr(args, "pretrained_model_name", "none") + pretrained_model_path = getattr(args, "pretrained_model_path", "") + if pretrained_model_name != "none" or pretrained_model_path: + self.load_state_dict( + load_pretrained_model( + pretrained_model_name, + pretrained_model_path, + ) + ) + if not getattr(args, "load_pretrained_model_output_layer", False): self.encoder.reset_output_layer_parameters() @staticmethod @@ -278,9 +285,12 @@ def base_architecture(args): @register_model_architecture("graphormer", "graphormer_base") def graphormer_base_architecture(args): - if args.pretrained_model_name == "pcqm4mv1_graphormer_base" or \ - args.pretrained_model_name == "pcqm4mv2_graphormer_base" or \ - args.pretrained_model_name == "pcqm4mv1_graphormer_base_for_molhiv": + pretrained_model_name = getattr(args, "pretrained_model_name", "none") + if pretrained_model_name in { + "pcqm4mv1_graphormer_base", + "pcqm4mv2_graphormer_base", + "pcqm4mv1_graphormer_base_for_molhiv", + }: args.encoder_layers = 12 args.encoder_attention_heads = 32 args.encoder_ffn_embed_dim = 768 diff --git a/graphormer/pretrain/__init__.py b/graphormer/pretrain/__init__.py index b473f36..6fd9921 100644 --- a/graphormer/pretrain/__init__.py +++ b/graphormer/pretrain/__init__.py @@ -1,19 +1,68 @@ -from torch.hub import load_state_dict_from_url +from collections.abc import Mapping +from pathlib import Path + +import torch import torch.distributed as dist +from torch.hub import load_state_dict_from_url PRETRAINED_MODEL_URLS = { - "pcqm4mv1_graphormer_base":"https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv1.pt", - "pcqm4mv2_graphormer_base":"https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv2.pt", - "oc20is2re_graphormer3d_base":"https://szheng.blob.core.windows.net/graphormer/modelzoo/oc20is2re/checkpoint_last_oc20_is2re.pt", # this pretrained model is temporarily unavailable - "pcqm4mv1_graphormer_base_for_molhiv":"https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_base_preln_pcqm4mv1_for_hiv.pt", + "pcqm4mv1_graphormer_base": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv1.pt", + "pcqm4mv2_graphormer_base": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv2.pt", + # This pretrained model is temporarily unavailable. + "oc20is2re_graphormer3d_base": "https://szheng.blob.core.windows.net/graphormer/modelzoo/oc20is2re/checkpoint_last_oc20_is2re.pt", + "pcqm4mv1_graphormer_base_for_molhiv": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_base_preln_pcqm4mv1_for_hiv.pt", } -def load_pretrained_model(pretrained_model_name): - if pretrained_model_name not in PRETRAINED_MODEL_URLS: - raise ValueError("Unknown pretrained model name %s", pretrained_model_name) - if not dist.is_initialized(): - return load_state_dict_from_url(PRETRAINED_MODEL_URLS[pretrained_model_name], progress=True)["model"] + +def _load_local_pretrained_model(pretrained_model_path): + checkpoint_path = Path(pretrained_model_path).expanduser() + if not checkpoint_path.exists(): + raise FileNotFoundError( + f"Pretrained model path does not exist: {checkpoint_path}" + ) + if not checkpoint_path.is_file(): + raise IsADirectoryError( + f"Pretrained model path is not a file: {checkpoint_path}" + ) + + checkpoint = torch.load(str(checkpoint_path), map_location="cpu") + if isinstance(checkpoint, Mapping): + model_state = checkpoint.get("model", checkpoint) else: - pretrained_model = load_state_dict_from_url(PRETRAINED_MODEL_URLS[pretrained_model_name], progress=True, file_name=f"{pretrained_model_name}_{dist.get_rank()}")["model"] + model_state = checkpoint + + if ( + not isinstance(model_state, Mapping) + or not model_state + or not all(isinstance(key, str) for key in model_state) + ): + raise ValueError( + "Local pretrained model must be a non-empty state dictionary or a " + "checkpoint containing a 'model' state dictionary" + ) + return model_state + + +def load_pretrained_model(pretrained_model_name="none", pretrained_model_path=None): + if pretrained_model_path: + if pretrained_model_name != "none": + raise ValueError( + "Set either pretrained_model_name or pretrained_model_path, not both" + ) + return _load_local_pretrained_model(pretrained_model_path) + + if pretrained_model_name not in PRETRAINED_MODEL_URLS: + raise ValueError(f"Unknown pretrained model name: {pretrained_model_name}") + if dist.is_initialized(): + checkpoint = load_state_dict_from_url( + PRETRAINED_MODEL_URLS[pretrained_model_name], + progress=True, + file_name=f"{pretrained_model_name}_{dist.get_rank()}", + ) dist.barrier() - return pretrained_model + else: + checkpoint = load_state_dict_from_url( + PRETRAINED_MODEL_URLS[pretrained_model_name], + progress=True, + ) + return checkpoint["model"] diff --git a/graphormer/tasks/graph_prediction.py b/graphormer/tasks/graph_prediction.py index 407b585..397275a 100644 --- a/graphormer/tasks/graph_prediction.py +++ b/graphormer/tasks/graph_prediction.py @@ -114,6 +114,11 @@ class GraphPredictionConfig(FairseqDataclass): metadata={"help": "name of used pretrained model"}, ) + pretrained_model_path: str = field( + default="", + metadata={"help": "path to a local pretrained model checkpoint"}, + ) + load_pretrained_model_output_layer: bool = field( default=False, metadata={"help": "whether to load the output layer of pretrained model"}, diff --git a/tests/test_pretrained_model.py b/tests/test_pretrained_model.py new file mode 100644 index 0000000..d7f09e2 --- /dev/null +++ b/tests/test_pretrained_model.py @@ -0,0 +1,107 @@ +import tempfile +import unittest +from pathlib import Path +from unittest import mock + +import torch + +from graphormer.pretrain import load_pretrained_model + + +class LoadPretrainedModelTest(unittest.TestCase): + def setUp(self): + self.temporary_directory = tempfile.TemporaryDirectory() + self.addCleanup(self.temporary_directory.cleanup) + self.directory = Path(self.temporary_directory.name) + self.model_state = {"encoder.weight": torch.tensor([1.0, 2.0])} + + def test_loads_fairseq_checkpoint(self): + checkpoint_path = self.directory / "checkpoint.pt" + torch.save({"model": self.model_state, "extra_state": {}}, checkpoint_path) + + loaded = load_pretrained_model( + pretrained_model_path=str(checkpoint_path) + ) + + torch.testing.assert_close( + loaded["encoder.weight"], + self.model_state["encoder.weight"], + ) + + def test_loads_raw_state_dictionary(self): + checkpoint_path = self.directory / "state_dict.pt" + torch.save(self.model_state, checkpoint_path) + + loaded = load_pretrained_model( + pretrained_model_path=str(checkpoint_path) + ) + + torch.testing.assert_close( + loaded["encoder.weight"], + self.model_state["encoder.weight"], + ) + + def test_rejects_missing_path(self): + checkpoint_path = self.directory / "missing.pt" + + with self.assertRaisesRegex(FileNotFoundError, "does not exist"): + load_pretrained_model(pretrained_model_path=str(checkpoint_path)) + + def test_rejects_directory_path(self): + with self.assertRaisesRegex(IsADirectoryError, "is not a file"): + load_pretrained_model(pretrained_model_path=str(self.directory)) + + def test_rejects_malformed_checkpoint(self): + checkpoint_path = self.directory / "malformed.pt" + torch.save(["not", "a", "state", "dictionary"], checkpoint_path) + + with self.assertRaisesRegex(ValueError, "state dictionary"): + load_pretrained_model(pretrained_model_path=str(checkpoint_path)) + + def test_rejects_name_and_path_together(self): + checkpoint_path = self.directory / "checkpoint.pt" + torch.save(self.model_state, checkpoint_path) + + with self.assertRaisesRegex(ValueError, "not both"): + load_pretrained_model( + pretrained_model_name="pcqm4mv2_graphormer_base", + pretrained_model_path=str(checkpoint_path), + ) + + def test_named_model_download_is_unchanged(self): + with mock.patch( + "graphormer.pretrain.dist.is_initialized", + return_value=False, + ), mock.patch( + "graphormer.pretrain.load_state_dict_from_url", + return_value={"model": self.model_state}, + ) as download: + loaded = load_pretrained_model("pcqm4mv2_graphormer_base") + + download.assert_called_once() + torch.testing.assert_close( + loaded["encoder.weight"], + self.model_state["encoder.weight"], + ) + + def test_local_model_loads_when_distributed_is_initialized(self): + checkpoint_path = self.directory / "checkpoint.pt" + torch.save({"model": self.model_state}, checkpoint_path) + + with mock.patch( + "graphormer.pretrain.dist.is_initialized", + return_value=True, + ), mock.patch("graphormer.pretrain.load_state_dict_from_url") as download: + loaded = load_pretrained_model( + pretrained_model_path=str(checkpoint_path) + ) + + download.assert_not_called() + torch.testing.assert_close( + loaded["encoder.weight"], + self.model_state["encoder.weight"], + ) + + +if __name__ == "__main__": + unittest.main()