AIDD / UniPath /remote /DiffCSP-official /diffcsp /script_utils.py
Wthinker's picture
Publish AIDD open-source resources
4947683 verified
Raw History Blame Contribute Delete
4.61 kB
import chemparse
import numpy as np
import torch
from torch.utils.data import Dataset
from torch_geometric.data import Data
chemical_symbols = [
# 0
'X',
# 1
'H', 'He',
# 2
'Li', 'Be', 'B', 'C', 'N', 'O', 'F', 'Ne',
# 3
'Na', 'Mg', 'Al', 'Si', 'P', 'S', 'Cl', 'Ar',
# 4
'K', 'Ca', 'Sc', 'Ti', 'V', 'Cr', 'Mn', 'Fe', 'Co', 'Ni', 'Cu', 'Zn',
'Ga', 'Ge', 'As', 'Se', 'Br', 'Kr',
# 5
'Rb', 'Sr', 'Y', 'Zr', 'Nb', 'Mo', 'Tc', 'Ru', 'Rh', 'Pd', 'Ag', 'Cd',
'In', 'Sn', 'Sb', 'Te', 'I', 'Xe',
# 6
'Cs', 'Ba', 'La', 'Ce', 'Pr', 'Nd', 'Pm', 'Sm', 'Eu', 'Gd', 'Tb', 'Dy',
'Ho', 'Er', 'Tm', 'Yb', 'Lu',
'Hf', 'Ta', 'W', 'Re', 'Os', 'Ir', 'Pt', 'Au', 'Hg', 'Tl', 'Pb', 'Bi',
'Po', 'At', 'Rn',
# 7
'Fr', 'Ra', 'Ac', 'Th', 'Pa', 'U', 'Np', 'Pu', 'Am', 'Cm', 'Bk',
'Cf', 'Es', 'Fm', 'Md', 'No', 'Lr',
'Rf', 'Db', 'Sg', 'Bh', 'Hs', 'Mt', 'Ds', 'Rg', 'Cn', 'Nh', 'Fl', 'Mc',
'Lv', 'Ts', 'Og']
class SampleDataset(Dataset):
def __init__(self, formula, num_evals):
super().__init__()
self.formula = formula
self.num_evals = num_evals
self.get_structure()
def get_structure(self):
self.composition = chemparse.parse_formula(self.formula)
chem_list = []
for elem in self.composition:
num_int = int(self.composition[elem])
chem_list.extend([chemical_symbols.index(elem)] * num_int)
self.chem_list = chem_list
def __len__(self) -> int:
return self.num_evals
def __getitem__(self, index):
return Data(
atom_types=torch.LongTensor(self.chem_list),
num_atoms=len(self.chem_list),
num_nodes=len(self.chem_list),
)
train_dist = {
'perov_5': [0, 0, 0, 0, 0, 1],
'carbon_24': [0.0,
0.0,
0.0,
0.0,
0.0,
0.0,
0.3250697750779839,
0.0,
0.27795107535708424,
0.0,
0.15383352487276308,
0.0,
0.11246100804465604,
0.0,
0.04958134953209654,
0.0,
0.038745690362830404,
0.0,
0.019044491873255624,
0.0,
0.010178952552946971,
0.0,
0.007059596125430964,
0.0,
0.006074536200952225],
'perov': [0, 0, 0, 0, 0, 1],
'carbon': [0.0,
0.0,
0.0,
0.0,
0.0,
0.0,
0.3250697750779839,
0.0,
0.27795107535708424,
0.0,
0.15383352487276308,
0.0,
0.11246100804465604,
0.0,
0.04958134953209654,
0.0,
0.038745690362830404,
0.0,
0.019044491873255624,
0.0,
0.010178952552946971,
0.0,
0.007059596125430964,
0.0,
0.006074536200952225],
'mp_20' : [0.0,
0.0021742334905660377,
0.021079009433962265,
0.019826061320754717,
0.15271226415094338,
0.047132959905660375,
0.08464770047169812,
0.021079009433962265,
0.07808814858490566,
0.03434551886792453,
0.0972877358490566,
0.013303360849056603,
0.09669811320754718,
0.02155807783018868,
0.06522700471698113,
0.014372051886792452,
0.06703272405660378,
0.00972877358490566,
0.053176591981132074,
0.010576356132075472,
0.08995430424528301]
}
class GenDataset(Dataset):
def __init__(self, dataset, total_num):
super().__init__()
self.total_num = total_num
self.distribution = train_dist[dataset]
self.num_atoms = np.random.choice(len(self.distribution), total_num, p = self.distribution)
self.is_carbon = dataset == 'carbon_24'
def __len__(self) -> int:
return self.total_num
def __getitem__(self, index):
num_atom = self.num_atoms[index]
data = Data(
num_atoms=torch.LongTensor([num_atom]),
num_nodes=num_atom,
)
if self.is_carbon:
data.atom_types = torch.LongTensor([6] * num_atom)
return data