Download UniPath/remote/DiffCSP-official/scripts/evaluate.py from BAAI/AIDD: direct link, hf CLI and curl.
- Browser
- Download file 3.63 kB
-
https://huggingface.co/BAAI/AIDD/resolve/main/UniPath/remote/DiffCSP-official/scripts/evaluate.py
- Command line
-
hf download hf://BAAI/AIDD/UniPath/remote/DiffCSP-official/scripts/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/BAAI/AIDD/resolve/main/UniPath/remote/DiffCSP-official/scripts/evaluate.py
3.63 kB
| import time | |
| import argparse | |
| import torch | |
| from tqdm import tqdm | |
| from torch.optim import Adam | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| from torch_geometric.data import Batch | |
| from diffcsp.eval_utils import load_model, lattices_to_params_shape, recommand_step_lr | |
| from pymatgen.core.structure import Structure | |
| from pymatgen.core.lattice import Lattice | |
| from pymatgen.symmetry.analyzer import SpacegroupAnalyzer | |
| from pyxtal.symmetry import Group | |
| import copy | |
| import numpy as np | |
| def diffusion(loader, model, num_evals, step_lr = 1e-5): | |
| frac_coords = [] | |
| num_atoms = [] | |
| atom_types = [] | |
| lattices = [] | |
| input_data_list = [] | |
| for idx, batch in enumerate(loader): | |
| if torch.cuda.is_available(): | |
| batch.cuda() | |
| batch_all_frac_coords = [] | |
| batch_all_lattices = [] | |
| batch_frac_coords, batch_num_atoms, batch_atom_types = [], [], [] | |
| batch_lattices = [] | |
| for eval_idx in range(num_evals): | |
| print(f'batch {idx} / {len(loader)}, sample {eval_idx} / {num_evals}') | |
| outputs, traj = model.sample(batch, step_lr = step_lr) | |
| batch_frac_coords.append(outputs['frac_coords'].detach().cpu()) | |
| batch_num_atoms.append(outputs['num_atoms'].detach().cpu()) | |
| batch_atom_types.append(outputs['atom_types'].detach().cpu()) | |
| batch_lattices.append(outputs['lattices'].detach().cpu()) | |
| frac_coords.append(torch.stack(batch_frac_coords, dim=0)) | |
| num_atoms.append(torch.stack(batch_num_atoms, dim=0)) | |
| atom_types.append(torch.stack(batch_atom_types, dim=0)) | |
| lattices.append(torch.stack(batch_lattices, dim=0)) | |
| input_data_list = input_data_list + batch.to_data_list() | |
| frac_coords = torch.cat(frac_coords, dim=1) | |
| num_atoms = torch.cat(num_atoms, dim=1) | |
| atom_types = torch.cat(atom_types, dim=1) | |
| lattices = torch.cat(lattices, dim=1) | |
| lengths, angles = lattices_to_params_shape(lattices) | |
| input_data_batch = Batch.from_data_list(input_data_list) | |
| return ( | |
| frac_coords, atom_types, lattices, lengths, angles, num_atoms, input_data_batch | |
| ) | |
| def main(args): | |
| # load_data if do reconstruction. | |
| model_path = Path(args.model_path) | |
| model, test_loader, cfg = load_model( | |
| model_path, | |
| load_data=True, | |
| ) | |
| if torch.cuda.is_available(): | |
| model.to('cuda') | |
| print('Evaluate the diffusion model.') | |
| step_lr = args.step_lr if args.step_lr >= 0 else recommand_step_lr['csp' if args.num_evals == 1 else 'csp_multi'][args.dataset] | |
| start_time = time.time() | |
| (frac_coords, atom_types, lattices, lengths, angles, num_atoms, input_data_batch) = diffusion( | |
| test_loader, model, args.num_evals, step_lr) | |
| if args.label == '': | |
| diff_out_name = 'eval_diff.pt' | |
| else: | |
| diff_out_name = f'eval_diff_{args.label}.pt' | |
| torch.save({ | |
| 'eval_setting': args, | |
| 'input_data_batch': input_data_batch, | |
| 'frac_coords': frac_coords, | |
| 'num_atoms': num_atoms, | |
| 'atom_types': atom_types, | |
| 'lattices': lattices, | |
| 'lengths': lengths, | |
| 'angles': angles, | |
| 'time': time.time() - start_time, | |
| }, model_path / diff_out_name) | |
| if __name__ == '__main__': | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--model_path', required=True) | |
| parser.add_argument('--dataset', required=True) | |
| parser.add_argument('--step_lr', default=-1, type=float) | |
| parser.add_argument('--num_evals', default=1, type=int) | |
| parser.add_argument('--label', default='') | |
| args = parser.parse_args() | |
| main(args) | |