Wthinker's picture
Publish AIDD open-source resources
4947683 verified
Raw History Blame Contribute Delete
3.53 kB
import time
import argparse
from diffcsp.script_utils import SampleDataset
import torch
from tqdm import tqdm
from torch.optim import Adam
from pathlib import Path
from types import SimpleNamespace
from torch_geometric.loader import DataLoader
from diffcsp.eval_utils import load_model, lattices_to_params_shape, get_crystals_list
from pymatgen.core.structure import Structure
from pymatgen.core.lattice import Lattice
from pymatgen.symmetry.analyzer import SpacegroupAnalyzer
from pymatgen.io.cif import CifWriter
from pyxtal.symmetry import Group
from p_tqdm import p_map
import os
def diffusion(loader, model, step_lr):
frac_coords = []
num_atoms = []
atom_types = []
lattices = []
input_data_list = []
for idx, batch in enumerate(loader):
if torch.cuda.is_available():
batch.cuda()
outputs, traj = model.sample(batch, step_lr = step_lr)
frac_coords.append(outputs['frac_coords'].detach().cpu())
num_atoms.append(outputs['num_atoms'].detach().cpu())
atom_types.append(outputs['atom_types'].detach().cpu())
lattices.append(outputs['lattices'].detach().cpu())
frac_coords = torch.cat(frac_coords, dim=0)
num_atoms = torch.cat(num_atoms, dim=0)
atom_types = torch.cat(atom_types, dim=0)
lattices = torch.cat(lattices, dim=0)
lengths, angles = lattices_to_params_shape(lattices)
return (
frac_coords, atom_types, lattices, lengths, angles, num_atoms
)
def get_pymatgen(crystal_array):
frac_coords = crystal_array['frac_coords']
atom_types = crystal_array['atom_types']
lengths = crystal_array['lengths']
angles = crystal_array['angles']
try:
structure = Structure(
lattice=Lattice.from_parameters(
*(lengths.tolist() + angles.tolist())),
species=atom_types, coords=frac_coords, coords_are_cartesian=False)
return structure
except:
return None
def main(args):
# load_data if do reconstruction.
model_path = Path(args.model_path)
model, _, cfg = load_model(
model_path, load_data=False)
if torch.cuda.is_available():
model.to('cuda')
tar_dir = os.path.join(args.save_path, args.formula)
os.makedirs(tar_dir, exist_ok=True)
print('Evaluate the diffusion model.')
test_set = SampleDataset(args.formula, args.num_evals)
test_loader = DataLoader(test_set, batch_size = min(args.batch_size, args.num_evals))
start_time = time.time()
(frac_coords, atom_types, lattices, lengths, angles, num_atoms) = diffusion(test_loader, model, args.step_lr)
crystal_list = get_crystals_list(frac_coords, atom_types, lengths, angles, num_atoms)
strcuture_list = p_map(get_pymatgen, crystal_list)
for i,structure in enumerate(strcuture_list):
tar_file = os.path.join(tar_dir, f"{args.formula}_{i+1}.cif")
if structure is not None:
writer = CifWriter(structure)
writer.write_file(tar_file)
else:
print(f"{i+1} Error Structure.")
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--model_path', required=True)
parser.add_argument('--save_path', required=True)
parser.add_argument('--formula', required=True)
parser.add_argument('--num_evals', default=1, type=int)
parser.add_argument('--batch_size', default=500, type=int)
parser.add_argument('--step_lr', default=1e-5, type=float)
args = parser.parse_args()
main(args)