Download UniPath/remote/DiffCSP-official/diffcsp/pl_data/dataset.py from BAAI/AIDD: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/BAAI/AIDD/resolve/main/UniPath/remote/DiffCSP-official/diffcsp/pl_data/dataset.py
- Command line
-
hf download hf://BAAI/AIDD/UniPath/remote/DiffCSP-official/diffcsp/pl_data/dataset.py
-
curl -L -o dataset.py https://huggingface.co/BAAI/AIDD/resolve/main/UniPath/remote/DiffCSP-official/diffcsp/pl_data/dataset.py
10.6 kB
| import hydra | |
| import omegaconf | |
| import torch | |
| import pandas as pd | |
| from omegaconf import ValueNode, OmegaConf | |
| from torch.utils.data import Dataset | |
| import os | |
| from torch_geometric.data import Data | |
| import pickle | |
| import numpy as np | |
| from diffcsp.common.utils import PROJECT_ROOT | |
| from diffcsp.common.data_utils import ( | |
| preprocess, preprocess_tensors, add_scaled_lattice_prop) | |
| class CrystDataset(Dataset): | |
| def __init__(self, name: ValueNode, path: ValueNode, | |
| prop: ValueNode, niggli: ValueNode, primitive: ValueNode, | |
| graph_method: ValueNode, preprocess_workers: ValueNode, | |
| lattice_scale_method: ValueNode, save_path: ValueNode, tolerance: ValueNode, use_space_group: ValueNode, use_pos_index: ValueNode, | |
| require_order: ValueNode = True, | |
| **kwargs): | |
| super().__init__() | |
| self.path = path | |
| self.name = name | |
| if os.path.exists(path): | |
| self.df = pd.read_csv(path) | |
| self.prop = prop | |
| if not isinstance(self.prop, str): | |
| self.prop = OmegaConf.to_container(self.prop) | |
| self.require_order = require_order | |
| self.niggli = niggli | |
| self.primitive = primitive | |
| self.graph_method = graph_method | |
| self.lattice_scale_method = lattice_scale_method | |
| self.use_space_group = use_space_group | |
| self.use_pos_index = use_pos_index | |
| self.tolerance = tolerance | |
| self.preprocess(save_path, preprocess_workers, prop) | |
| add_scaled_lattice_prop(self.cached_data, lattice_scale_method) | |
| self.lattice_scaler = None | |
| self.scaler = None | |
| def preprocess(self, save_path, preprocess_workers, prop): | |
| if os.path.exists(save_path): | |
| self.cached_data = torch.load(save_path) | |
| else: | |
| cached_data = preprocess( | |
| self.path, | |
| preprocess_workers, | |
| niggli=self.niggli, | |
| primitive=self.primitive, | |
| graph_method=self.graph_method, | |
| prop_list=[prop] if isinstance(prop, str) else prop, | |
| use_space_group=self.use_space_group, | |
| tol=self.tolerance, | |
| require_order=self.require_order) | |
| torch.save(cached_data, save_path) | |
| self.cached_data = cached_data | |
| def __len__(self) -> int: | |
| return len(self.cached_data) | |
| def __getitem__(self, index): | |
| data_dict = self.cached_data[index] | |
| # scaler is set in DataModule set stage | |
| if isinstance(self.prop, str): | |
| prop = {"y": self.scaler.transform(data_dict[self.prop]).view(1, -1)} | |
| elif isinstance(self.prop, list): | |
| prop = {} | |
| for prop_name in self.prop: | |
| prop_val = data_dict[prop_name] | |
| if isinstance(prop_val, str): | |
| prop_val = eval(prop_val) | |
| prop[prop_name] = torch.tensor(prop_val, dtype = torch.float32) | |
| if prop_name in ["energy", "stress"]: | |
| prop[prop_name] = prop[prop_name].unsqueeze(0) | |
| else: | |
| raise ValueError(f"prop should be string or list of string, but get {type(prop)}") | |
| # 解包 graph_arrays,处理可能缺少 edge_indices 和 to_jimages 的情况 | |
| graph_arrays = data_dict['graph_arrays'] | |
| n_elements = len(graph_arrays) | |
| if n_elements == 7: | |
| # 正常情况,包含所有元素 | |
| (frac_coords, atom_types, lengths, angles, | |
| edge_indices, to_jimages, num_atoms) = graph_arrays | |
| elif n_elements == 5: | |
| # 缺少 edge_indices 和 to_jimages,设为空数组 | |
| (frac_coords, atom_types, lengths, angles, | |
| num_atoms) = graph_arrays | |
| edge_indices = np.empty((0, 2), dtype=np.int64) # 空的边索引 | |
| to_jimages = np.empty((0, 3), dtype=np.int64) # 空的晶像偏移 | |
| else: | |
| raise ValueError(f"Unexpected number of elements in graph_arrays: {n_elements}") | |
| # atom_coords are fractional coordinates | |
| # edge_index is incremented during batching | |
| # https://pytorch-geometric.readthedocs.io/en/latest/notes/batching.html | |
| data = Data( | |
| frac_coords=torch.Tensor(frac_coords), | |
| atom_types=torch.LongTensor(atom_types), | |
| lengths=torch.Tensor(lengths).view(1, -1), | |
| angles=torch.Tensor(angles).view(1, -1), | |
| edge_index=torch.LongTensor( | |
| edge_indices.T).contiguous(), # shape (2, num_edges) | |
| to_jimages=torch.LongTensor(to_jimages), | |
| num_atoms=num_atoms, | |
| num_bonds=edge_indices.shape[0], | |
| num_nodes=num_atoms, # special attribute used for batching in pytorch geometric | |
| **prop | |
| ) | |
| if "t_flow" in data_dict: | |
| data.t_flow = torch.tensor(data_dict["t_flow"], dtype = torch.float32) | |
| if "graph_arrays_stable" in data_dict: | |
| (frac_coords_stable, atom_types_stable, lengths_stable, angles_stable, num_atoms_stable) = data_dict['graph_arrays_stable'] | |
| data.frac_coords_stable = torch.Tensor(frac_coords_stable) | |
| data.atom_types_stable = torch.LongTensor(atom_types_stable) | |
| data.lengths_stable = torch.Tensor(lengths_stable).view(1, -1) | |
| data.angles_stable = torch.Tensor(angles_stable).view(1, -1) | |
| data.num_atoms_stable = num_atoms_stable | |
| data.num_nodes_stable = num_atoms_stable | |
| if "graph_arrays_initial" in data_dict: | |
| (frac_coords_initial, atom_types_initial, lengths_initial, angles_initial, edge_indices_initial, | |
| to_jimages_initial, num_atoms_initial) = data_dict['graph_arrays_initial'] | |
| data.frac_coords_initial = torch.Tensor(frac_coords_initial) | |
| data.atom_types_initial = torch.LongTensor(atom_types_initial) | |
| data.lengths_initial = torch.Tensor(lengths_initial).view(1, -1) | |
| data.angles_initial = torch.Tensor(angles_initial).view(1, -1) | |
| data.edge_index_initial = torch.LongTensor( | |
| edge_indices_initial.T).contiguous() | |
| data.to_jimages_initial = torch.LongTensor(to_jimages_initial) | |
| data.num_atoms_initial = num_atoms_initial | |
| data.num_bonds_initial = edge_indices_initial.shape[0] | |
| data.num_nodes_initial = num_atoms_initial | |
| if self.use_space_group: | |
| data.spacegroup = torch.LongTensor([data_dict['spacegroup']]) | |
| data.ops = torch.Tensor(data_dict['wyckoff_ops']) | |
| data.anchor_index = torch.LongTensor(data_dict['anchors']) | |
| if self.use_pos_index: | |
| pos_dic = {} | |
| indexes = [] | |
| for atom in atom_types: | |
| pos_dic[atom] = pos_dic.get(atom, 0) + 1 | |
| indexes.append(pos_dic[atom] - 1) | |
| data.index = torch.LongTensor(indexes) | |
| return data | |
| def __repr__(self) -> str: | |
| return f"CrystDataset({self.name=}, {self.path=})" | |
| class TensorCrystDataset(Dataset): | |
| def __init__(self, crystal_array_list, niggli, primitive, | |
| graph_method, preprocess_workers, | |
| lattice_scale_method, **kwargs): | |
| super().__init__() | |
| self.niggli = niggli | |
| self.primitive = primitive | |
| self.graph_method = graph_method | |
| self.lattice_scale_method = lattice_scale_method | |
| self.cached_data = preprocess_tensors( | |
| crystal_array_list, | |
| niggli=self.niggli, | |
| primitive=self.primitive, | |
| graph_method=self.graph_method) | |
| add_scaled_lattice_prop(self.cached_data, lattice_scale_method) | |
| self.lattice_scaler = None | |
| self.scaler = None | |
| def __len__(self) -> int: | |
| return len(self.cached_data) | |
| def __getitem__(self, index): | |
| data_dict = self.cached_data[index] | |
| (frac_coords, atom_types, lengths, angles, edge_indices, | |
| to_jimages, num_atoms) = data_dict['graph_arrays'] | |
| # atom_coords are fractional coordinates | |
| # edge_index is incremented during batching | |
| # https://pytorch-geometric.readthedocs.io/en/latest/notes/batching.html | |
| data = Data( | |
| frac_coords=torch.Tensor(frac_coords), | |
| atom_types=torch.LongTensor(atom_types), | |
| lengths=torch.Tensor(lengths).view(1, -1), | |
| angles=torch.Tensor(angles).view(1, -1), | |
| edge_index=torch.LongTensor( | |
| edge_indices.T).contiguous(), # shape (2, num_edges) | |
| to_jimages=torch.LongTensor(to_jimages), | |
| num_atoms=num_atoms, | |
| num_bonds=edge_indices.shape[0], | |
| num_nodes=num_atoms, # special attribute used for batching in pytorch geometric | |
| ) | |
| if "graph_arrays_initial" in data_dict: | |
| (frac_coords_initial, atom_types_initial, lengths_initial, angles_initial, edge_indices_initial, | |
| to_jimages_initial, num_atoms_initial) = data_dict['graph_arrays_initial'] | |
| data.frac_coords_initial = torch.Tensor(frac_coords_initial) | |
| data.atom_types_initial = torch.LongTensor(atom_types_initial) | |
| data.lengths_initial = torch.Tensor(lengths_initial).view(1, -1) | |
| data.angles_initial = torch.Tensor(angles_initial).view(1, -1) | |
| data.edge_index_initial = torch.LongTensor( | |
| edge_indices_initial.T).contiguous() | |
| data.to_jimages_initial = torch.LongTensor(to_jimages_initial) | |
| data.num_atoms_initial = num_atoms_initial | |
| data.num_bonds_initial = edge_indices_initial.shape[0] | |
| data.num_nodes_initial = num_atoms_initial | |
| return data | |
| def __repr__(self) -> str: | |
| return f"TensorCrystDataset(len: {len(self.cached_data)})" | |
| def main(cfg: omegaconf.DictConfig): | |
| from torch_geometric.data import Batch | |
| from diffcsp.common.data_utils import get_scaler_from_data_list | |
| dataset: CrystDataset = hydra.utils.instantiate( | |
| cfg.data.datamodule.datasets.train, _recursive_=False | |
| ) | |
| lattice_scaler = get_scaler_from_data_list( | |
| dataset.cached_data, | |
| key='scaled_lattice') | |
| scaler = get_scaler_from_data_list( | |
| dataset.cached_data, | |
| key=dataset.prop) | |
| dataset.lattice_scaler = lattice_scaler | |
| dataset.scaler = scaler | |
| data_list = [dataset[i] for i in range(len(dataset))] | |
| batch = Batch.from_data_list(data_list) | |
| return batch | |
| if __name__ == "__main__": | |
| main() | |