| import numpy as np |
| import pandas as pd |
| import networkx as nx |
| import torch |
| import copy |
| import itertools |
|
|
| from pymatgen.core.structure import Structure |
| from pymatgen.core.lattice import Lattice |
| from pymatgen.analysis.graphs import StructureGraph |
| from pymatgen.analysis import local_env |
|
|
| from networkx.algorithms.components import is_connected |
|
|
| from sklearn.metrics import accuracy_score, recall_score, precision_score |
|
|
| from torch_scatter import scatter |
|
|
| from p_tqdm import p_umap |
|
|
|
|
| |
| |
| OFFSET_LIST = [ |
| [-1, -1, -1], |
| [-1, -1, 0], |
| [-1, -1, 1], |
| [-1, 0, -1], |
| [-1, 0, 0], |
| [-1, 0, 1], |
| [-1, 1, -1], |
| [-1, 1, 0], |
| [-1, 1, 1], |
| [0, -1, -1], |
| [0, -1, 0], |
| [0, -1, 1], |
| [0, 0, -1], |
| [0, 0, 0], |
| [0, 0, 1], |
| [0, 1, -1], |
| [0, 1, 0], |
| [0, 1, 1], |
| [1, -1, -1], |
| [1, -1, 0], |
| [1, -1, 1], |
| [1, 0, -1], |
| [1, 0, 0], |
| [1, 0, 1], |
| [1, 1, -1], |
| [1, 1, 0], |
| [1, 1, 1], |
| ] |
|
|
| EPSILON = 1e-5 |
|
|
| chemical_symbols = [ |
| |
| 'X', |
| |
| 'H', 'He', |
| |
| 'Li', 'Be', 'B', 'C', 'N', 'O', 'F', 'Ne', |
| |
| 'Na', 'Mg', 'Al', 'Si', 'P', 'S', 'Cl', 'Ar', |
| |
| 'K', 'Ca', 'Sc', 'Ti', 'V', 'Cr', 'Mn', 'Fe', 'Co', 'Ni', 'Cu', 'Zn', |
| 'Ga', 'Ge', 'As', 'Se', 'Br', 'Kr', |
| |
| 'Rb', 'Sr', 'Y', 'Zr', 'Nb', 'Mo', 'Tc', 'Ru', 'Rh', 'Pd', 'Ag', 'Cd', |
| 'In', 'Sn', 'Sb', 'Te', 'I', 'Xe', |
| |
| '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', |
| |
| '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'] |
|
|
|
|
| CrystalNN = local_env.CrystalNN( |
| distance_cutoffs=None, x_diff_weight=-1, porous_adjustment=False) |
|
|
|
|
| def build_crystal(crystal_str, niggli=True, primitive=False): |
| """Build crystal from cif string.""" |
| crystal = Structure.from_str(crystal_str, fmt='cif') |
|
|
| if primitive: |
| crystal = crystal.get_primitive_structure() |
|
|
| if niggli: |
| crystal = crystal.get_reduced_structure() |
|
|
| canonical_crystal = Structure( |
| lattice=Lattice.from_parameters(*crystal.lattice.parameters), |
| species=crystal.species, |
| coords=crystal.frac_coords, |
| coords_are_cartesian=False, |
| ) |
| |
| |
| return canonical_crystal |
|
|
|
|
| def build_crystal_graph(crystal, graph_method='crystalnn'): |
| """ |
| """ |
|
|
| if graph_method == 'crystalnn': |
| crystal_graph = StructureGraph.with_local_env_strategy( |
| crystal, CrystalNN) |
| elif graph_method == 'none': |
| pass |
| else: |
| raise NotImplementedError |
|
|
| frac_coords = crystal.frac_coords |
| atom_types = crystal.atomic_numbers |
| lattice_parameters = crystal.lattice.parameters |
| lengths = lattice_parameters[:3] |
| angles = lattice_parameters[3:] |
|
|
| assert np.allclose(crystal.lattice.matrix, |
| lattice_params_to_matrix(*lengths, *angles)) |
|
|
| edge_indices, to_jimages = [], [] |
| if graph_method != 'none': |
| for i, j, to_jimage in crystal_graph.graph.edges(data='to_jimage'): |
| edge_indices.append([j, i]) |
| to_jimages.append(to_jimage) |
| edge_indices.append([i, j]) |
| to_jimages.append(tuple(-tj for tj in to_jimage)) |
|
|
| atom_types = np.array(atom_types) |
| lengths, angles = np.array(lengths), np.array(angles) |
| edge_indices = np.array(edge_indices) |
| to_jimages = np.array(to_jimages) |
| num_atoms = atom_types.shape[0] |
|
|
| return frac_coords, atom_types, lengths, angles, edge_indices, to_jimages, num_atoms |
|
|
|
|
| def abs_cap(val, max_abs_val=1): |
| """ |
| Returns the value with its absolute value capped at max_abs_val. |
| Particularly useful in passing values to trignometric functions where |
| numerical errors may result in an argument > 1 being passed in. |
| https://github.com/materialsproject/pymatgen/blob/b789d74639aa851d7e5ee427a765d9fd5a8d1079/pymatgen/util/num.py#L15 |
| Args: |
| val (float): Input value. |
| max_abs_val (float): The maximum absolute value for val. Defaults to 1. |
| Returns: |
| val if abs(val) < 1 else sign of val * max_abs_val. |
| """ |
| return max(min(val, max_abs_val), -max_abs_val) |
|
|
|
|
| def lattice_params_to_matrix(a, b, c, alpha, beta, gamma): |
| """Converts lattice from abc, angles to matrix. |
| https://github.com/materialsproject/pymatgen/blob/b789d74639aa851d7e5ee427a765d9fd5a8d1079/pymatgen/core/lattice.py#L311 |
| """ |
| angles_r = np.radians([alpha, beta, gamma]) |
| cos_alpha, cos_beta, cos_gamma = np.cos(angles_r) |
| sin_alpha, sin_beta, sin_gamma = np.sin(angles_r) |
|
|
| val = (cos_alpha * cos_beta - cos_gamma) / (sin_alpha * sin_beta) |
| |
| val = abs_cap(val) |
| gamma_star = np.arccos(val) |
|
|
| vector_a = [a * sin_beta, 0.0, a * cos_beta] |
| vector_b = [ |
| -b * sin_alpha * np.cos(gamma_star), |
| b * sin_alpha * np.sin(gamma_star), |
| b * cos_alpha, |
| ] |
| vector_c = [0.0, 0.0, float(c)] |
| return np.array([vector_a, vector_b, vector_c]) |
|
|
|
|
| def lattice_params_to_matrix_torch(lengths, angles): |
| """Batched torch version to compute lattice matrix from params. |
| |
| lengths: torch.Tensor of shape (N, 3), unit A |
| angles: torch.Tensor of shape (N, 3), unit degree |
| """ |
| angles_r = torch.deg2rad(angles) |
| coses = torch.cos(angles_r) |
| sins = torch.sin(angles_r) |
|
|
| val = (coses[:, 0] * coses[:, 1] - coses[:, 2]) / (sins[:, 0] * sins[:, 1]) |
| |
| val = torch.clamp(val, -1., 1.) |
| gamma_star = torch.arccos(val) |
|
|
| vector_a = torch.stack([ |
| lengths[:, 0] * sins[:, 1], |
| torch.zeros(lengths.size(0), device=lengths.device), |
| lengths[:, 0] * coses[:, 1]], dim=1) |
| vector_b = torch.stack([ |
| -lengths[:, 1] * sins[:, 0] * torch.cos(gamma_star), |
| lengths[:, 1] * sins[:, 0] * torch.sin(gamma_star), |
| lengths[:, 1] * coses[:, 0]], dim=1) |
| vector_c = torch.stack([ |
| torch.zeros(lengths.size(0), device=lengths.device), |
| torch.zeros(lengths.size(0), device=lengths.device), |
| lengths[:, 2]], dim=1) |
|
|
| return torch.stack([vector_a, vector_b, vector_c], dim=1) |
|
|
|
|
| def compute_volume(batch_lattice): |
| """Compute volume from batched lattice matrix |
| |
| batch_lattice: (N, 3, 3) |
| """ |
| vector_a, vector_b, vector_c = torch.unbind(batch_lattice, dim=1) |
| return torch.abs(torch.einsum('bi,bi->b', vector_a, |
| torch.cross(vector_b, vector_c, dim=1))) |
|
|
|
|
| def lengths_angles_to_volume(lengths, angles): |
| lattice = lattice_params_to_matrix_torch(lengths, angles) |
| return compute_volume(lattice) |
|
|
|
|
| def lattice_matrix_to_params(matrix): |
| lengths = np.sqrt(np.sum(matrix ** 2, axis=1)).tolist() |
|
|
| angles = np.zeros(3) |
| for i in range(3): |
| j = (i + 1) % 3 |
| k = (i + 2) % 3 |
| angles[i] = abs_cap(np.dot(matrix[j], matrix[k]) / |
| (lengths[j] * lengths[k])) |
| angles = np.arccos(angles) * 180.0 / np.pi |
| a, b, c = lengths |
| alpha, beta, gamma = angles |
| return a, b, c, alpha, beta, gamma |
|
|
|
|
| def frac_to_cart_coords( |
| frac_coords, |
| lengths, |
| angles, |
| num_atoms, |
| ): |
| lattice = lattice_params_to_matrix_torch(lengths, angles) |
| lattice_nodes = torch.repeat_interleave(lattice, num_atoms, dim=0) |
| pos = torch.einsum('bi,bij->bj', frac_coords, lattice_nodes) |
|
|
| return pos |
|
|
|
|
| def cart_to_frac_coords( |
| cart_coords, |
| lengths, |
| angles, |
| num_atoms, |
| ): |
| lattice = lattice_params_to_matrix_torch(lengths, angles) |
| |
| inv_lattice = torch.linalg.pinv(lattice) |
| inv_lattice_nodes = torch.repeat_interleave(inv_lattice, num_atoms, dim=0) |
| frac_coords = torch.einsum('bi,bij->bj', cart_coords, inv_lattice_nodes) |
| return (frac_coords % 1.) |
|
|
|
|
| def get_pbc_distances( |
| coords, |
| edge_index, |
| lengths, |
| angles, |
| to_jimages, |
| num_atoms, |
| num_bonds, |
| coord_is_cart=False, |
| return_offsets=False, |
| return_distance_vec=False, |
| ): |
| lattice = lattice_params_to_matrix_torch(lengths, angles) |
|
|
| if coord_is_cart: |
| pos = coords |
| else: |
| lattice_nodes = torch.repeat_interleave(lattice, num_atoms, dim=0) |
| pos = torch.einsum('bi,bij->bj', coords, lattice_nodes) |
|
|
| j_index, i_index = edge_index |
|
|
| distance_vectors = pos[j_index] - pos[i_index] |
|
|
| |
| lattice_edges = torch.repeat_interleave(lattice, num_bonds, dim=0) |
| offsets = torch.einsum('bi,bij->bj', to_jimages.float(), lattice_edges) |
| distance_vectors += offsets |
|
|
| |
| distances = distance_vectors.norm(dim=-1) |
|
|
| out = { |
| "edge_index": edge_index, |
| "distances": distances, |
| } |
|
|
| if return_distance_vec: |
| out["distance_vec"] = distance_vectors |
|
|
| if return_offsets: |
| out["offsets"] = offsets |
|
|
| return out |
|
|
|
|
| def radius_graph_pbc_wrapper(data, radius, max_num_neighbors_threshold, device): |
| cart_coords = frac_to_cart_coords( |
| data.frac_coords, data.lengths, data.angles, data.num_atoms) |
| return radius_graph_pbc( |
| cart_coords, data.lengths, data.angles, data.num_atoms, radius, |
| max_num_neighbors_threshold, device) |
|
|
|
|
| def radius_graph_pbc(cart_coords, lengths, angles, num_atoms, |
| radius, max_num_neighbors_threshold, device, |
| topk_per_pair=None): |
| """Computes pbc graph edges under pbc. |
| |
| topk_per_pair: (num_atom_pairs,), select topk edges per atom pair |
| |
| Note: topk should take into account self-self edge for (i, i) |
| """ |
| batch_size = len(num_atoms) |
|
|
| |
| atom_pos = cart_coords |
|
|
| |
| num_atoms_per_image = num_atoms |
| num_atoms_per_image_sqr = (num_atoms_per_image ** 2).long() |
|
|
| |
| index_offset = ( |
| torch.cumsum(num_atoms_per_image, dim=0) - num_atoms_per_image |
| ) |
|
|
| index_offset_expand = torch.repeat_interleave( |
| index_offset, num_atoms_per_image_sqr |
| ) |
| num_atoms_per_image_expand = torch.repeat_interleave( |
| num_atoms_per_image, num_atoms_per_image_sqr |
| ) |
|
|
| |
| |
| |
| |
| |
| num_atom_pairs = torch.sum(num_atoms_per_image_sqr) |
| index_sqr_offset = ( |
| torch.cumsum(num_atoms_per_image_sqr, dim=0) - num_atoms_per_image_sqr |
| ) |
| index_sqr_offset = torch.repeat_interleave( |
| index_sqr_offset, num_atoms_per_image_sqr |
| ) |
| atom_count_sqr = ( |
| torch.arange(num_atom_pairs, device=device) - index_sqr_offset |
| ) |
|
|
| |
| |
| index1 = ( |
| (atom_count_sqr // num_atoms_per_image_expand) |
| ).long() + index_offset_expand |
| index2 = ( |
| atom_count_sqr % num_atoms_per_image_expand |
| ).long() + index_offset_expand |
| |
| pos1 = torch.index_select(atom_pos, 0, index1) |
| pos2 = torch.index_select(atom_pos, 0, index2) |
|
|
| unit_cell = torch.tensor(OFFSET_LIST, device=device).float() |
| num_cells = len(unit_cell) |
| unit_cell_per_atom = unit_cell.view(1, num_cells, 3).repeat( |
| len(index2), 1, 1 |
| ) |
| unit_cell = torch.transpose(unit_cell, 0, 1) |
| unit_cell_batch = unit_cell.view(1, 3, num_cells).expand( |
| batch_size, -1, -1 |
| ) |
|
|
| |
| lattice = lattice_params_to_matrix_torch(lengths, angles) |
|
|
| |
| data_cell = torch.transpose(lattice, 1, 2) |
| pbc_offsets = torch.bmm(data_cell, unit_cell_batch) |
| pbc_offsets_per_atom = torch.repeat_interleave( |
| pbc_offsets, num_atoms_per_image_sqr, dim=0 |
| ) |
|
|
| |
| pos1 = pos1.view(-1, 3, 1).expand(-1, -1, num_cells) |
| pos2 = pos2.view(-1, 3, 1).expand(-1, -1, num_cells) |
| index1 = index1.view(-1, 1).repeat(1, num_cells).view(-1) |
| index2 = index2.view(-1, 1).repeat(1, num_cells).view(-1) |
| |
| pos2 = pos2 + pbc_offsets_per_atom |
|
|
| |
| atom_distance_sqr = torch.sum((pos1 - pos2) ** 2, dim=1) |
|
|
| if topk_per_pair is not None: |
| assert topk_per_pair.size(0) == num_atom_pairs |
| atom_distance_sqr_sort_index = torch.argsort(atom_distance_sqr, dim=1) |
| assert atom_distance_sqr_sort_index.size() == (num_atom_pairs, num_cells) |
| atom_distance_sqr_sort_index = ( |
| atom_distance_sqr_sort_index + |
| torch.arange(num_atom_pairs, device=device)[:, None] * num_cells).view(-1) |
| topk_mask = (torch.arange(num_cells, device=device)[None, :] < |
| topk_per_pair[:, None]) |
| topk_mask = topk_mask.view(-1) |
| topk_indices = atom_distance_sqr_sort_index.masked_select(topk_mask) |
|
|
| topk_mask = torch.zeros(num_atom_pairs * num_cells, device=device) |
| topk_mask.scatter_(0, topk_indices, 1.) |
| topk_mask = topk_mask.bool() |
|
|
| atom_distance_sqr = atom_distance_sqr.view(-1) |
|
|
| |
| mask_within_radius = torch.le(atom_distance_sqr, radius * radius) |
| |
| mask_not_same = torch.gt(atom_distance_sqr, 0.0001) |
| mask = torch.logical_and(mask_within_radius, mask_not_same) |
| index1 = torch.masked_select(index1, mask) |
| index2 = torch.masked_select(index2, mask) |
| unit_cell = torch.masked_select( |
| unit_cell_per_atom.view(-1, 3), mask.view(-1, 1).expand(-1, 3) |
| ) |
| unit_cell = unit_cell.view(-1, 3) |
| if topk_per_pair is not None: |
| topk_mask = torch.masked_select(topk_mask, mask) |
|
|
| num_neighbors = torch.zeros(len(cart_coords), device=device) |
| num_neighbors.index_add_(0, index1, torch.ones(len(index1), device=device)) |
| num_neighbors = num_neighbors.long() |
| max_num_neighbors = torch.max(num_neighbors).long() |
|
|
| |
| _max_neighbors = copy.deepcopy(num_neighbors) |
| _max_neighbors[ |
| _max_neighbors > max_num_neighbors_threshold |
| ] = max_num_neighbors_threshold |
| _num_neighbors = torch.zeros(len(cart_coords) + 1, device=device).long() |
| _natoms = torch.zeros(num_atoms.shape[0] + 1, device=device).long() |
| _num_neighbors[1:] = torch.cumsum(_max_neighbors, dim=0) |
| _natoms[1:] = torch.cumsum(num_atoms, dim=0) |
| num_neighbors_image = ( |
| _num_neighbors[_natoms[1:]] - _num_neighbors[_natoms[:-1]] |
| ) |
|
|
| |
| if ( |
| max_num_neighbors <= max_num_neighbors_threshold |
| or max_num_neighbors_threshold <= 0 |
| ): |
| if topk_per_pair is None: |
| return torch.stack((index2, index1)), unit_cell, num_neighbors_image |
| else: |
| return torch.stack((index2, index1)), unit_cell, num_neighbors_image, topk_mask |
|
|
| atom_distance_sqr = torch.masked_select(atom_distance_sqr, mask) |
|
|
| |
| |
| distance_sort = torch.zeros( |
| len(cart_coords) * max_num_neighbors, device=device |
| ).fill_(radius * radius + 1.0) |
|
|
| |
| index_neighbor_offset = torch.cumsum(num_neighbors, dim=0) - num_neighbors |
| index_neighbor_offset_expand = torch.repeat_interleave( |
| index_neighbor_offset, num_neighbors |
| ) |
| index_sort_map = ( |
| index1 * max_num_neighbors |
| + torch.arange(len(index1), device=device) |
| - index_neighbor_offset_expand |
| ) |
| distance_sort.index_copy_(0, index_sort_map, atom_distance_sqr) |
| distance_sort = distance_sort.view(len(cart_coords), max_num_neighbors) |
|
|
| |
| distance_sort, index_sort = torch.sort(distance_sort, dim=1) |
| |
| distance_sort = distance_sort[:, :max_num_neighbors_threshold] |
| index_sort = index_sort[:, :max_num_neighbors_threshold] |
|
|
| |
| index_sort = index_sort + index_neighbor_offset.view(-1, 1).expand( |
| -1, max_num_neighbors_threshold |
| ) |
| |
| mask_within_radius = torch.le(distance_sort, radius * radius) |
| index_sort = torch.masked_select(index_sort, mask_within_radius) |
|
|
| |
| |
| mask_num_neighbors = torch.zeros(len(index1), device=device).bool() |
| mask_num_neighbors.index_fill_(0, index_sort, True) |
|
|
| |
| index1 = torch.masked_select(index1, mask_num_neighbors) |
| index2 = torch.masked_select(index2, mask_num_neighbors) |
| unit_cell = torch.masked_select( |
| unit_cell.view(-1, 3), mask_num_neighbors.view(-1, 1).expand(-1, 3) |
| ) |
| unit_cell = unit_cell.view(-1, 3) |
|
|
| if topk_per_pair is not None: |
| topk_mask = torch.masked_select(topk_mask, mask_num_neighbors) |
|
|
| edge_index = torch.stack((index2, index1)) |
|
|
| if topk_per_pair is None: |
| return edge_index, unit_cell, num_neighbors_image |
| else: |
| return edge_index, unit_cell, num_neighbors_image, topk_mask |
|
|
|
|
| def min_distance_sqr_pbc(cart_coords1, cart_coords2, lengths, angles, |
| num_atoms, device, return_vector=False, |
| return_to_jimages=False): |
| """Compute the pbc distance between atoms in cart_coords1 and cart_coords2. |
| This function assumes that cart_coords1 and cart_coords2 have the same number of atoms |
| in each data point. |
| returns: |
| basic return: |
| min_atom_distance_sqr: (N_atoms, ) |
| return_vector == True: |
| min_atom_distance_vector: vector pointing from cart_coords1 to cart_coords2, (N_atoms, 3) |
| return_to_jimages == True: |
| to_jimages: (N_atoms, 3), position of cart_coord2 relative to cart_coord1 in pbc |
| """ |
| batch_size = len(num_atoms) |
|
|
| |
| pos1 = cart_coords1 |
| pos2 = cart_coords2 |
|
|
| unit_cell = torch.tensor(OFFSET_LIST, device=device).float() |
| num_cells = len(unit_cell) |
| unit_cell_per_atom = unit_cell.view(1, num_cells, 3).repeat( |
| len(cart_coords2), 1, 1 |
| ) |
| unit_cell = torch.transpose(unit_cell, 0, 1) |
| unit_cell_batch = unit_cell.view(1, 3, num_cells).expand( |
| batch_size, -1, -1 |
| ) |
|
|
| |
| lattice = lattice_params_to_matrix_torch(lengths, angles) |
|
|
| |
| data_cell = torch.transpose(lattice, 1, 2) |
| pbc_offsets = torch.bmm(data_cell, unit_cell_batch) |
| pbc_offsets_per_atom = torch.repeat_interleave( |
| pbc_offsets, num_atoms, dim=0 |
| ) |
|
|
| |
| pos1 = pos1.view(-1, 3, 1).expand(-1, -1, num_cells) |
| pos2 = pos2.view(-1, 3, 1).expand(-1, -1, num_cells) |
| |
| pos2 = pos2 + pbc_offsets_per_atom |
|
|
| |
| |
| atom_distance_vector = pos1 - pos2 |
| atom_distance_sqr = torch.sum(atom_distance_vector ** 2, dim=1) |
|
|
| min_atom_distance_sqr, min_indices = atom_distance_sqr.min(dim=-1) |
|
|
| return_list = [min_atom_distance_sqr] |
|
|
| if return_vector: |
| min_indices = min_indices[:, None, None].repeat([1, 3, 1]) |
|
|
| min_atom_distance_vector = torch.gather( |
| atom_distance_vector, 2, min_indices).squeeze(-1) |
|
|
| return_list.append(min_atom_distance_vector) |
|
|
| if return_to_jimages: |
| to_jimages = unit_cell.T[min_indices].long() |
| return_list.append(to_jimages) |
|
|
| return return_list[0] if len(return_list) == 1 else return_list |
|
|
|
|
| class StandardScalerTorch(object): |
| """Normalizes the targets of a dataset.""" |
|
|
| def __init__(self, means=None, stds=None): |
| self.means = means |
| self.stds = stds |
|
|
| def fit(self, X): |
| X = torch.tensor(X, dtype=torch.float) |
| self.means = torch.mean(X, dim=0) |
| |
| self.stds = torch.std(X, dim=0, unbiased=False) + EPSILON |
|
|
| def transform(self, X): |
| X = torch.tensor(X, dtype=torch.float) |
| return (X - self.means) / self.stds |
|
|
| def inverse_transform(self, X): |
| X = torch.tensor(X, dtype=torch.float) |
| return X * self.stds + self.means |
|
|
| def match_device(self, tensor): |
| if self.means.device != tensor.device: |
| self.means = self.means.to(tensor.device) |
| self.stds = self.stds.to(tensor.device) |
|
|
| def copy(self): |
| return StandardScalerTorch( |
| means=self.means.clone().detach(), |
| stds=self.stds.clone().detach()) |
|
|
| def __repr__(self) -> str: |
| return ( |
| f"{self.__class__.__name__}(" |
| f"means: {self.means.tolist()}, " |
| f"stds: {self.stds.tolist()})" |
| ) |
|
|
|
|
| def get_scaler_from_data_list(data_list, key): |
| targets_list = [d[key] for d in data_list] |
| if isinstance(targets_list[0], torch.Tensor): |
| targets_list = [t.numpy() for t in targets_list] |
| targets = torch.tensor(targets_list) |
| scaler = StandardScalerTorch() |
| scaler.fit(targets) |
| return scaler |
|
|
|
|
| def preprocess(input_file, num_workers, niggli, primitive, graph_method, |
| prop_list): |
| df = pd.read_pickle(input_file) |
|
|
| def process_one(row, niggli, primitive, graph_method, prop_list): |
| crystal_str = row['cif'] |
| crystal = build_crystal( |
| crystal_str, niggli=niggli, primitive=primitive) |
| graph_arrays = build_crystal_graph(crystal, graph_method) |
| properties = {k: row[k] for k in prop_list if k in row.keys()} |
| result_dict = { |
| 'mp_id': row['material_id'], |
| 'cif': crystal_str, |
| 'graph_arrays': graph_arrays, |
| 'spacegroup.number': row['spacegroup.number'], |
| 'pretty_formula': row['pretty_formula'], |
| } |
| result_dict.update(properties) |
| return result_dict |
|
|
| unordered_results = p_umap( |
| process_one, |
| [df.iloc[idx] for idx in range(len(df))], |
| [niggli] * len(df), |
| [primitive] * len(df), |
| [graph_method] * len(df), |
| [prop_list] * len(df), |
| num_cpus=num_workers) |
|
|
| mpid_to_results = {result['mp_id']: result for result in unordered_results} |
| ordered_results = [mpid_to_results[df.iloc[idx]['material_id']] |
| for idx in range(len(df))] |
|
|
| return ordered_results |
|
|
|
|
| def preprocess_tensors(crystal_array_list, niggli, primitive, graph_method): |
| def process_one(batch_idx, crystal_array, niggli, primitive, graph_method): |
| frac_coords = crystal_array['frac_coords'] |
| atom_types = crystal_array['atom_types'] |
| lengths = crystal_array['lengths'] |
| angles = crystal_array['angles'] |
| crystal = Structure( |
| lattice=Lattice.from_parameters( |
| *(lengths.tolist() + angles.tolist())), |
| species=atom_types, |
| coords=frac_coords, |
| coords_are_cartesian=False) |
| graph_arrays = build_crystal_graph(crystal, graph_method) |
| result_dict = { |
| 'batch_idx': batch_idx, |
| 'graph_arrays': graph_arrays, |
| } |
| return result_dict |
|
|
| unordered_results = p_umap( |
| process_one, |
| list(range(len(crystal_array_list))), |
| crystal_array_list, |
| [niggli] * len(crystal_array_list), |
| [primitive] * len(crystal_array_list), |
| [graph_method] * len(crystal_array_list), |
| num_cpus=30, |
| ) |
| ordered_results = list( |
| sorted(unordered_results, key=lambda x: x['batch_idx'])) |
| return ordered_results |
|
|
|
|
| def add_scaled_lattice_prop(data_list, lattice_scale_method): |
| for dict in data_list: |
| graph_arrays = dict['graph_arrays'] |
| |
| lengths = graph_arrays[2] |
| angles = graph_arrays[3] |
| num_atoms = graph_arrays[-1] |
| assert lengths.shape[0] == angles.shape[0] == 3 |
| assert isinstance(num_atoms, int) |
|
|
| if lattice_scale_method == 'scale_length': |
| lengths = lengths / float(num_atoms)**(1/3) |
|
|
| dict['scaled_lattice'] = np.concatenate([lengths, angles]) |
|
|
|
|
| def mard(targets, preds): |
| """Mean absolute relative difference.""" |
| assert torch.all(targets > 0.) |
| return torch.mean(torch.abs(targets - preds) / targets) |
|
|
|
|
| def batch_accuracy_precision_recall( |
| pred_edge_probs, |
| edge_overlap_mask, |
| num_bonds |
| ): |
| if (pred_edge_probs is None and edge_overlap_mask is None and |
| num_bonds is None): |
| return 0., 0., 0. |
| pred_edges = pred_edge_probs.max(dim=1)[1].float() |
| target_edges = edge_overlap_mask.float() |
|
|
| start_idx = 0 |
| accuracies, precisions, recalls = [], [], [] |
| for num_bond in num_bonds.tolist(): |
| pred_edge = pred_edges.narrow( |
| 0, start_idx, num_bond).detach().cpu().numpy() |
| target_edge = target_edges.narrow( |
| 0, start_idx, num_bond).detach().cpu().numpy() |
|
|
| accuracies.append(accuracy_score(target_edge, pred_edge)) |
| precisions.append(precision_score( |
| target_edge, pred_edge, average='binary')) |
| recalls.append(recall_score(target_edge, pred_edge, average='binary')) |
|
|
| start_idx = start_idx + num_bond |
|
|
| return np.mean(accuracies), np.mean(precisions), np.mean(recalls) |
|
|
|
|
| class StandardScaler: |
| """A :class:`StandardScaler` normalizes the features of a dataset. |
| When it is fit on a dataset, the :class:`StandardScaler` learns the |
| mean and standard deviation across the 0th axis. |
| When transforming a dataset, the :class:`StandardScaler` subtracts the |
| means and divides by the standard deviations. |
| """ |
|
|
| def __init__(self, means=None, stds=None, replace_nan_token=None): |
| """ |
| :param means: An optional 1D numpy array of precomputed means. |
| :param stds: An optional 1D numpy array of precomputed standard deviations. |
| :param replace_nan_token: A token to use to replace NaN entries in the features. |
| """ |
| self.means = means |
| self.stds = stds |
| self.replace_nan_token = replace_nan_token |
|
|
| def fit(self, X): |
| """ |
| Learns means and standard deviations across the 0th axis of the data :code:`X`. |
| :param X: A list of lists of floats (or None). |
| :return: The fitted :class:`StandardScaler` (self). |
| """ |
| X = np.array(X).astype(float) |
| self.means = np.nanmean(X, axis=0) |
| self.stds = np.nanstd(X, axis=0) |
| self.means = np.where(np.isnan(self.means), |
| np.zeros(self.means.shape), self.means) |
| self.stds = np.where(np.isnan(self.stds), |
| np.ones(self.stds.shape), self.stds) |
| self.stds = np.where(self.stds == 0, np.ones( |
| self.stds.shape), self.stds) |
|
|
| return self |
|
|
| def transform(self, X): |
| """ |
| Transforms the data by subtracting the means and dividing by the standard deviations. |
| :param X: A list of lists of floats (or None). |
| :return: The transformed data with NaNs replaced by :code:`self.replace_nan_token`. |
| """ |
| X = np.array(X).astype(float) |
| transformed_with_nan = (X - self.means) / self.stds |
| transformed_with_none = np.where( |
| np.isnan(transformed_with_nan), self.replace_nan_token, transformed_with_nan) |
|
|
| return transformed_with_none |
|
|
| def inverse_transform(self, X): |
| """ |
| Performs the inverse transformation by multiplying by the standard deviations and adding the means. |
| :param X: A list of lists of floats. |
| :return: The inverse transformed data with NaNs replaced by :code:`self.replace_nan_token`. |
| """ |
| X = np.array(X).astype(float) |
| transformed_with_nan = X * self.stds + self.means |
| transformed_with_none = np.where( |
| np.isnan(transformed_with_nan), self.replace_nan_token, transformed_with_nan) |
|
|
| return transformed_with_none |
|
|