Skip to content
Merged
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
14 changes: 14 additions & 0 deletions examples/property_prediction/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 (%)|
--------------|---------|------------|
Expand Down
34 changes: 21 additions & 13 deletions graphormer/evaluate/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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)
Expand All @@ -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):
Expand Down
22 changes: 16 additions & 6 deletions graphormer/models/graphormer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
73 changes: 61 additions & 12 deletions graphormer/pretrain/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
5 changes: 5 additions & 0 deletions graphormer/tasks/graph_prediction.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"},
Expand Down
107 changes: 107 additions & 0 deletions tests/test_pretrained_model.py
Original file line number Diff line number Diff line change
@@ -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()
Loading