From b4ba938f2252445beea67867817e66a1739e3432 Mon Sep 17 00:00:00 2001 From: Victor Schmidt Date: Wed, 29 Nov 2023 16:26:55 -0500 Subject: [PATCH] import changes from #48 --- configs/models/tasks/is2re.yaml | 2 + ocpmodels/datasets/lmdb_dataset.py | 199 +++++++++++++++++++---------- ocpmodels/trainers/base_trainer.py | 29 ++++- 3 files changed, 154 insertions(+), 76 deletions(-) diff --git a/configs/models/tasks/is2re.yaml b/configs/models/tasks/is2re.yaml index cf47f159de..787e20295f 100644 --- a/configs/models/tasks/is2re.yaml +++ b/configs/models/tasks/is2re.yaml @@ -16,6 +16,8 @@ default: otf_graph: False max_num_neighbors: 40 mode: train + adsorbates: all # {"*O", "*OH", "*OH2", "*H"} + adsorbates_ref_dir: /network/scratch/s/schmidtv/ocp/datasets/ocp/per_ads dataset: default_val: val_id train: diff --git a/ocpmodels/datasets/lmdb_dataset.py b/ocpmodels/datasets/lmdb_dataset.py index ec953b4a28..0a7abaea85 100644 --- a/ocpmodels/datasets/lmdb_dataset.py +++ b/ocpmodels/datasets/lmdb_dataset.py @@ -6,6 +6,7 @@ """ import bisect +import json import logging import pickle import time @@ -16,7 +17,7 @@ import numpy as np import torch from torch.utils.data import Dataset -from torch_geometric.data import Batch, HeteroData, Data +from torch_geometric.data import Batch from ocpmodels.common.registry import registry from ocpmodels.common.utils import pyg2_data_transform @@ -36,15 +37,35 @@ class LmdbDataset(Dataset): config (dict): Dataset configuration transform (callable, optional): Data transform function. (default: :obj:`None`) + fa_frames (str, optional): type of frame averaging method applied, if any. + adsorbates (str, optional): comma-separated list of adsorbates to filter. + If None or "all", no filtering is applied. + (default: None) + adsorbates_ref_dir: where metadata files for adsorbates are stored. + (default: "/network/scratch/s/schmidtv/ocp/datasets/ocp/per_ads") """ - def __init__(self, config, transform=None, fa_frames=None): - super(LmdbDataset, self).__init__() + def __init__( + self, + config, + transform=None, + fa_frames=None, + lmdb_glob=None, + adsorbates=None, + adsorbates_ref_dir=None, + ): + super().__init__() self.config = config + self.adsorbates = adsorbates + self.adsorbates_ref_dir = adsorbates_ref_dir self.path = Path(self.config["src"]) if not self.path.is_file(): db_paths = sorted(self.path.glob("*.lmdb")) + if lmdb_glob: + db_paths = [ + p for p in db_paths if any(lg in p.stem for lg in lmdb_glob) + ] assert len(db_paths) > 0, f"No LMDBs found in '{self.path}'" self.metadata_path = self.path / "metadata.npz" @@ -58,7 +79,7 @@ def __init__(self, config, transform=None, fa_frames=None): else: length = self.envs[-1].stat()["entries"] assert length is not None, f"Could not find length of LMDB {db_path}" - self._keys.append(list(range(length))) + self._keys.append([str(i).encode("ascii") for i in range(length)]) keylens = [len(k) for k in self._keys] self._keylen_cumulative = np.cumsum(keylens).tolist() @@ -71,14 +92,78 @@ def __init__(self, config, transform=None, fa_frames=None): ] self.num_samples = len(self._keys) + self.filter_per_adsorbates() self.transform = transform - self.fa_frames = fa_frames + self.fa_method = fa_frames + + def filter_per_adsorbates(self): + """Filter the dataset to only include structures with a specific + adsorbate. + """ + # no adsorbates specified, or asked for all: return + if not self.adsorbates or self.adsorbates == "all": + return + + # val_ood_ads and val_ood_both don't have targeted adsorbates + if self.config["src"].split("/")[-2] in {"val_ood_ads", "val_ood_both"}: + return + + # make set of adsorbates from a list or a string. If a string, split on comma. + ads = [] + if isinstance(self.adsorbates, str): + if "," in self.adsorbates: + ads = [a.strip() for a in self.adsorbates.split(",")] + else: + ads = [self.adsorbates] + else: + ads = self.adsorbates + ads = set(ads) + + # find reference file for this dataset + ref_path = self.adsorbates_ref_dir + if not ref_path: + print("No adsorbate reference directory provided as `adsorbate_ref_dir`.") + return + ref_path = Path(ref_path) + if not ref_path.is_dir(): + print(f"Adsorbate reference directory {ref_path} does not exist.") + return + pattern = "-".join(self.path.parts[-3:]) + candidates = list(ref_path.glob(f"*{pattern}*.json")) + if not candidates: + print(f"No adsorbate reference files found for {self.path.name}.") + return + if len(candidates) > 1: + print( + f"Multiple adsorbate reference files found for {self.path.name}." + "Using the first one." + ) + ref = json.loads(candidates[0].read_text()) + + # find dataset indices with the appropriate adsorbates + allowed_idxs = set( + str(i).encode("ascii") + for i, a in zip(ref["ds_idx"], ref["ads_symbols"]) + if a in ads + ) + + # filter the dataset indices + if isinstance(self._keys[0], bytes): + self._keys = [i for i in self._keys if i in allowed_idxs] + self.num_samples = len(self._keys) + else: + assert isinstance(self._keys[0], list) + self._keys = [[i for i in k if i in allowed_idxs] for k in self._keys] + keylens = [len(k) for k in self._keys] + self._keylen_cumulative = np.cumsum(keylens).tolist() + self.num_samples = sum(keylens) + + assert self.num_samples > 0, f"No samples found for adsorbates {ads}." def __len__(self): return self.num_samples - def __getitem__(self, idx): - t0 = time.time_ns() + def get_pickled_from_db(self, idx): if not self.path.is_file(): # Figure out which db this should be indexed from. db_idx = bisect.bisect(self._keylen_cumulative, idx) @@ -89,16 +174,20 @@ def __getitem__(self, idx): assert el_idx >= 0 # Return features. - datapoint_pickled = ( - self.envs[db_idx] - .begin() - .get(f"{self._keys[db_idx][el_idx]}".encode("ascii")) + return ( + f"{db_idx}_{el_idx}", + self.envs[db_idx].begin().get(self._keys[db_idx][el_idx]), ) - data_object = pyg2_data_transform(pickle.loads(datapoint_pickled)) - data_object.id = f"{db_idx}_{el_idx}" - else: - datapoint_pickled = self.env.begin().get(self._keys[idx]) - data_object = pyg2_data_transform(pickle.loads(datapoint_pickled)) + + return None, self.env.begin().get(self._keys[idx]) + + def __getitem__(self, idx): + t0 = time.time_ns() + + el_id, datapoint_pickled = self.get_pickled_from_db(idx) + data_object = pyg2_data_transform(pickle.loads(datapoint_pickled)) + if el_id: + data_object.id = el_id t1 = time.time_ns() if self.transform is not None: @@ -112,6 +201,7 @@ def __getitem__(self, idx): data_object.load_time = load_time data_object.transform_time = transform_time data_object.total_get_time = total_get_time + data_object.idx_in_dataset = idx return data_object @@ -137,6 +227,27 @@ def close_db(self): self.env.close() +@registry.register_dataset("deup_lmdb") +class DeupDataset(LmdbDataset): + def __init__(self, all_datasets_configs, deup_split, transform=None): + super().__init__( + all_datasets_configs[deup_split], + lmdb_glob=deup_split.replace("deup-", "").split("-"), + ) + ocp_splits = deup_split.split("-")[1:] + self.ocp_datasets = { + d: LmdbDataset(all_datasets_configs[d], transform) for d in ocp_splits + } + + def __getitem__(self, idx): + _, datapoint_pickled = self.get_pickled_from_db(idx) + deup_sample = pickle.loads(datapoint_pickled) + ocp_sample = self.ocp_datasets[deup_sample["ds"]][deup_sample["idx_in_dataset"]] + for k, v in deup_sample.items(): + setattr(ocp_sample, f"deup_{k}", v) + return ocp_sample + + class SinglePointLmdbDataset(LmdbDataset): def __init__(self, config, transform=None): super(SinglePointLmdbDataset, self).__init__(config, transform) @@ -157,24 +268,8 @@ def __init__(self, config, transform=None): ) -# In this function, we combine a list of samples into a batch. Notice that we first create the batch, then we fix -# the neighbor problem: that some elements in the batch don't have edges, which pytorch geometric doesn't handle well -# and which leads to errors in the forward step. -def data_list_collater(data_list, otf_graph=False): # Check if len(batch) is ever used - # FIRST, MAKE BATCH - - if ( # This is for indfaenet - type(data_list[0]) is tuple and type(data_list[0][0]) is Data - ): - adsorbates = [system[0] for system in data_list] - catalysts = [system[1] for system in data_list] - - ads_batch = Batch.from_data_list(adsorbates) - cat_batch = Batch.from_data_list(catalysts) - else: - batch = Batch.from_data_list(data_list) - - # THEN, FIX NEIGHBOR PROBLEM +def data_list_collater(data_list, otf_graph=False): + batch = Batch.from_data_list(data_list) if ( not otf_graph @@ -192,40 +287,4 @@ def data_list_collater(data_list, otf_graph=False): # Check if len(batch) is ev "LMDB does not contain edge index information, set otf_graph=True" ) - elif ( # This is for indfaenet - not otf_graph and type(data_list[0]) is tuple and type(data_list[0][0]) is Data - ): - batches = [ads_batch, cat_batch] - lists = [adsorbates, catalysts] - for batch, list_type in zip(batches, lists): - n_neighbors = [] - for i, data in enumerate(list_type): - n_index = data.edge_index[1, :] - n_neighbors.append(n_index.shape[0]) - batch.neighbors = torch.tensor(n_neighbors) - - return batches - - elif not otf_graph and type(data_list[0]) is HeteroData: # This is for afaenet - # First, fix the neighborhood dimension. - n_neighbors_ads = [] - n_neighbors_cat = [] - for i, data in enumerate(data_list): - n_index_ads = data["adsorbate", "is_close", "adsorbate"].edge_index - n_index_cat = data["catalyst", "is_close", "catalyst"].edge_index - n_neighbors_ads.append(n_index_ads[1, :].shape[0]) - n_neighbors_cat.append(n_index_cat[1, :].shape[0]) - batch["adsorbate"].neighbors = torch.tensor(n_neighbors_ads) - batch["catalyst"].neighbors = torch.tensor(n_neighbors_cat) - - # Then, fix the edge index between ads and cats. - sender, receiver = batch["is_disc"].edge_index - ads_to_cat = torch.stack([sender, receiver + batch["adsorbate"].num_nodes]) - cat_to_ads = torch.stack([ads_to_cat[1], ads_to_cat[0]]) - batch["is_disc"].edge_index = torch.concat([ads_to_cat, cat_to_ads], dim=1) - - batch["is_disc"].edge_weight = torch.concat( - [batch["is_disc"].edge_weight, -batch["is_disc"].edge_weight], dim=0 - ) - return batch diff --git a/ocpmodels/trainers/base_trainer.py b/ocpmodels/trainers/base_trainer.py index 17c2243f0b..0b36b08405 100644 --- a/ocpmodels/trainers/base_trainer.py +++ b/ocpmodels/trainers/base_trainer.py @@ -8,12 +8,13 @@ import errno import logging import os +import pickle import random import time -import pickle from abc import ABC, abstractmethod from collections import defaultdict from copy import deepcopy +from uuid import uuid4 import numpy as np import torch @@ -26,7 +27,7 @@ from torch.utils.data import DataLoader, Subset from torch_geometric.data import Batch from tqdm import tqdm -from uuid import uuid4 + from ocpmodels.common import dist_utils from ocpmodels.common.data_parallel import ( BalancedBatchSampler, @@ -36,7 +37,12 @@ from ocpmodels.common.graph_transforms import RandomReflect, RandomRotate from ocpmodels.common.registry import registry from ocpmodels.common.timer import Times -from ocpmodels.common.utils import JOB_ID, get_commit_hash, save_checkpoint, resolve +from ocpmodels.common.utils import ( + JOB_ID, + get_commit_hash, + resolve, + save_checkpoint, +) from ocpmodels.datasets.data_transforms import FrameAveraging, get_transforms from ocpmodels.modules.evaluator import Evaluator from ocpmodels.modules.exponential_moving_average import ( @@ -292,18 +298,29 @@ def load_datasets(self): if self.data_mode == "separate": self.datasets[split] = registry.get_dataset_class("separate")( - ds_conf, transform=transform + ds_conf, + transform=transform, + adsorbates=self.config.get("adsorbates"), + adsorbates_ref_dir=self.config.get("adsorbates_ref_dir"), ) elif self.data_mode == "heterogeneous": self.datasets[split] = registry.get_dataset_class("heterogeneous")( - ds_conf, transform=transform + ds_conf, + transform=transform, + adsorbates=self.config.get("adsorbates"), + adsorbates_ref_dir=self.config.get("adsorbates_ref_dir"), ) else: self.datasets[split] = registry.get_dataset_class( self.config["task"]["dataset"] - )(ds_conf, transform=transform) + )( + ds_conf, + transform=transform, + adsorbates=self.config.get("adsorbates"), + adsorbates_ref_dir=self.config.get("adsorbates_ref_dir"), + ) if self.config["lowest_energy_only"]: with open(