Wthinker's picture
Publish AIDD open-source resources
4947683 verified
Raw History Blame Contribute Delete
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)})"
@hydra.main(config_path=str(PROJECT_ROOT / "conf"), config_name="default", version_base="1.1")
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()