Spaces:
Running
Running
| """ | |
| This files includes data processing tools. | |
| """ | |
| import os | |
| import argparse | |
| import json | |
| from typing import Iterable, Literal | |
| import numpy as np | |
| import pandas as pd | |
| from sklearn.base import BaseEstimator, TransformerMixin | |
| from sklearn.preprocessing import StandardScaler | |
| from sklearn.feature_selection import VarianceThreshold | |
| from statsmodels.distributions.empirical_distribution import ECDF | |
| from datasets import load_dataset | |
| import torch | |
| from rdkit import Chem, DataStructs | |
| from rdkit.Chem import Descriptors, rdFingerprintGenerator, MACCSkeys | |
| from rdkit.Chem.rdchem import Mol | |
| from src.utils import ( | |
| TASKS, | |
| HF_TOKEN, | |
| USED_200_DESCR, | |
| Standardizer, | |
| load_pickle, | |
| write_pickle, | |
| KNOWN_DESCR, | |
| ) | |
| class SquashScaler(TransformerMixin, BaseEstimator): | |
| """ | |
| Scaler that performs sequential standardization, nonlinearity (tanh), and | |
| re-standardization. Inspired by DeepTox (Mayr et al., 2016) | |
| """ | |
| def __init__(self): | |
| self.scaler1 = StandardScaler() | |
| self.scaler2 = StandardScaler() | |
| def fit(self, X): | |
| _X = X.copy() | |
| _X = self.scaler1.fit_transform(_X) | |
| _X = np.tanh(_X) | |
| _X = self.scaler2.fit(_X) | |
| self.is_fitted_ = True | |
| return self | |
| def transform(self, X): | |
| _X = X.copy() | |
| _X = self.scaler1.transform(_X) | |
| _X = np.tanh(_X) | |
| return self.scaler2.transform(_X) | |
| def create_cleaned_mol_objects(smiles: list[str]) -> tuple[list[Mol], np.ndarray]: | |
| """This function creates cleaned RDKit mol objects from a list of SMILES. | |
| Args: | |
| smiles (list[str]): list of SMILES | |
| Returns: | |
| list[Mol]: list of cleaned molecules | |
| np.ndarray[bool]: mask that contains False at index `i`, if molecule in `smiles` at | |
| index `i` could not be cleaned and was removed. | |
| """ | |
| sm = Standardizer(canon_taut=True) | |
| clean_mol_mask = list() | |
| mols = list() | |
| for i, smile in enumerate(smiles): | |
| mol = Chem.MolFromSmiles(smile) | |
| standardized_mol, _ = sm.standardize_mol(mol) | |
| is_cleaned = standardized_mol is not None | |
| clean_mol_mask.append(is_cleaned) | |
| if not is_cleaned: | |
| continue | |
| can_mol = Chem.MolFromSmiles(Chem.MolToSmiles(standardized_mol)) | |
| mols.append(can_mol) | |
| return mols, np.array(clean_mol_mask) | |
| def create_ecfp_fps(mols: list[Mol], radius=None, fpsize=None) -> np.ndarray: | |
| """This function ECFP fingerprints for a list of molecules. | |
| Args: | |
| mols (list[Mol]): list of molecules | |
| Returns: | |
| np.ndarray: ECFP fingerprints of molecules | |
| """ | |
| ecfps = list() | |
| kwargs = {} | |
| if not fpsize is None: | |
| kwargs["fpSize"] = fpsize | |
| if not radius is None: | |
| kwargs["radius"] = radius | |
| for mol in mols: | |
| gen = rdFingerprintGenerator.GetMorganGenerator(countSimulation=True, **kwargs) | |
| fp_sparse_vec = gen.GetCountFingerprint(mol) | |
| fp = np.zeros((0,), np.int8) | |
| DataStructs.ConvertToNumpyArray(fp_sparse_vec, fp) | |
| ecfps.append(fp) | |
| return np.array(ecfps) | |
| def create_maccs_keys(mols: list[Mol]) -> np.ndarray: | |
| maccs = [MACCSkeys.GenMACCSKeys(x) for x in mols] | |
| return np.array(maccs) | |
| def get_tox_patterns(filepath: str): | |
| """This calculates tox features defined in tox_smarts.json. | |
| Args: | |
| mols: A list of Mol | |
| n_jobs: If >1 multiprocessing is used | |
| """ | |
| # load patterns | |
| with open(filepath) as f: | |
| smarts_list = [s[1] for s in json.load(f)] | |
| # Code does not work for this case | |
| assert len([s for s in smarts_list if ("AND" in s) and ("OR" in s)]) == 0 | |
| # Chem.MolFromSmarts takes a long time so it pays of to parse all the smarts first | |
| # and then use them for all molecules. This gives a huge speedup over existing code. | |
| # a list of patterns, whether to negate the match result and how to join them to obtain one boolean value | |
| all_patterns = [] | |
| for smarts in smarts_list: | |
| patterns = [] # list of smarts-patterns | |
| # value for each of the patterns above. Negates the values of the above later. | |
| negations = [] | |
| if " AND " in smarts: | |
| smarts = smarts.split(" AND ") | |
| merge_any = False # If an ' AND ' is found all 'subsmarts' have to match | |
| else: | |
| # If there is an ' OR ' present it's enough is any of the 'subsmarts' match. | |
| # This also accumulates smarts where neither ' OR ' nor ' AND ' occur | |
| smarts = smarts.split(" OR ") | |
| merge_any = True | |
| # for all subsmarts check if they are preceded by 'NOT ' | |
| for s in smarts: | |
| neg = s.startswith("NOT ") | |
| if neg: | |
| s = s[4:] | |
| patterns.append(Chem.MolFromSmarts(s)) | |
| negations.append(neg) | |
| all_patterns.append((patterns, negations, merge_any)) | |
| return all_patterns | |
| def create_tox_features(mols: list[Mol], patterns: list) -> np.ndarray: | |
| """Matches the tox patterns against a molecule. Returns a boolean array""" | |
| tox_data = [] | |
| for mol in mols: | |
| mol_features = [] | |
| for patts, negations, merge_any in patterns: | |
| matches = [mol.HasSubstructMatch(p) for p in patts] | |
| matches = [m != n for m, n in zip(matches, negations)] | |
| if merge_any: | |
| pres = any(matches) | |
| else: | |
| pres = all(matches) | |
| mol_features.append(pres) | |
| tox_data.append(np.array(mol_features)) | |
| return np.array(tox_data) | |
| def create_rdkit_descriptors(mols: list[Mol]) -> np.ndarray: | |
| """This function creates RDKit descriptors for a list of molecules. | |
| Args: | |
| mols (list[Mol]): list of molecules | |
| Returns: | |
| np.ndarray: RDKit descriptors of molecules | |
| """ | |
| rdkit_descriptors = list() | |
| for mol in mols: | |
| descrs = [] | |
| for _, descr_calc_fn in Descriptors._descList: | |
| descrs.append(descr_calc_fn(mol)) | |
| descrs = np.array(descrs) | |
| descrs = descrs[USED_200_DESCR] | |
| rdkit_descriptors.append(descrs) | |
| return np.array(rdkit_descriptors) | |
| def create_quantiles(raw_features: np.ndarray, ecdfs: list) -> np.ndarray: | |
| """Create quantile values for given features using the columns | |
| Args: | |
| raw_features (np.ndarray): values to put into quantiles | |
| ecdfs (list): ECDFs to use | |
| Returns: | |
| np.ndarray: computed quantiles | |
| """ | |
| quantiles = np.zeros_like(raw_features) | |
| for column in range(raw_features.shape[1]): | |
| raw_values = raw_features[:, column].reshape(-1) | |
| ecdf = ecdfs[column] | |
| q = ecdf(raw_values) | |
| quantiles[:, column] = q | |
| return quantiles | |
| def fill(features, mask, value=np.nan): | |
| n_mols = len(mask) | |
| n_features = features.shape[1] | |
| data = np.zeros(shape=(n_mols, n_features)) | |
| data.fill(value) | |
| data[~mask] = features | |
| return data | |
| def get_descriptor_dataset( | |
| data_path: str, | |
| descriptors: Iterable[str] | Literal["all"], | |
| scaler=None, | |
| save_scaler_path: str = "data/scaler.pkl", | |
| verbose=True, | |
| normalize: str = "standard", | |
| ): | |
| if descriptors == "all": | |
| descriptors = KNOWN_DESCR | |
| assert isinstance(descriptors, Iterable), "Passed descriptors are not iterable!" | |
| assert all( | |
| [descr in KNOWN_DESCR for descr in descriptors] | |
| ), f"Passed descriptors contains unknown descriptor types. Allowed descriptors: {KNOWN_DESCR}" | |
| print(f"Load: {data_path}") | |
| datafile = np.load(data_path) | |
| if not isinstance(datafile, np.ndarray): | |
| # concatenate all descriptors and normalize | |
| if "features" in datafile: | |
| data = datafile["features"] | |
| print("Features are already concatenated") | |
| else: | |
| data = np.concatenate([datafile[descr] for descr in descriptors], axis=1) | |
| print(f"Concatenated features with order: {descriptors}") | |
| labels = datafile["labels"] | |
| else: | |
| print("NPY file passed, cannot select specific descriptors") | |
| data, labels = datafile[:, :-12], datafile[:, -12:] | |
| if normalize != "none": | |
| data, scaler = normalize_features( | |
| data, | |
| scaler=scaler, | |
| save_scaler_path=save_scaler_path, | |
| verbose=verbose, | |
| normalization=normalize, | |
| ) | |
| # filter out unsanitized molecules | |
| mask = ~np.isnan(data).any(axis=1) | |
| data = data[mask] | |
| labels = labels[mask] | |
| assert data.shape[0] == labels.shape[0], ( | |
| f"Mismatch between data and labels: " | |
| f"data has {data.shape[0]} samples, but labels has {labels.shape[0]} samples." | |
| ) | |
| return (data, labels, scaler) | |
| def get_torch_descriptor_dataset( | |
| data_path: str, | |
| descriptors: list[str], | |
| scaler=None, | |
| save_scaler_path: str = "data/scaler.pkl", | |
| nan_to_num: int = -100, | |
| verbose=True, | |
| normalize: str = "standard", | |
| ) -> torch.utils.data.TensorDataset: | |
| data, labels, scaler = get_descriptor_dataset( | |
| data_path, | |
| descriptors, | |
| scaler, | |
| save_scaler_path, | |
| verbose=verbose, | |
| normalize=normalize, | |
| ) | |
| labels = np.nan_to_num(labels, nan=nan_to_num) | |
| dataset = torch.utils.data.TensorDataset( | |
| torch.FloatTensor(data), torch.LongTensor(labels) | |
| ) | |
| return dataset, scaler | |
| def get_tox21_split(token="", cvfold=None): | |
| """Retrieve Tox21 splits from HuggingFace with respect to given cvfold.""" | |
| ds = load_dataset("ml-jku/tox21", token=token) | |
| train_df = ds["train"].to_pandas() | |
| val_df = ds["validation"].to_pandas() | |
| if cvfold is None: | |
| return {"train": train_df, "validation": val_df} | |
| combined_df = pd.concat([train_df, val_df], ignore_index=True) | |
| cvfold = float(cvfold) | |
| # create new splits | |
| cvfold = float(cvfold) | |
| train_df = combined_df[combined_df.CVfold != cvfold] | |
| val_df = combined_df[combined_df.CVfold == cvfold] | |
| # exclude train mols that occur in the validation split | |
| val_inchikeys = set(val_df["inchikey"]) | |
| train_df = train_df[~train_df["inchikey"].isin(val_inchikeys)] | |
| return { | |
| "train": train_df.reset_index(drop=True), | |
| "validation": val_df.reset_index(drop=True), | |
| } | |
| def normalize_features( | |
| raw_features, | |
| scaler=None, | |
| save_scaler_path: str = "", | |
| verbose=True, | |
| normalization: str = "standard", | |
| ): | |
| if scaler is None: | |
| if normalization == "standard": | |
| scaler = StandardScaler() | |
| elif normalization == "squash": | |
| scaler = SquashScaler() | |
| scaler.fit(raw_features) | |
| if verbose: | |
| print("Fitted the StandardScaler") | |
| if save_scaler_path: | |
| write_pickle(save_scaler_path, scaler) | |
| if verbose: | |
| print(f"Saved the StandardScaler under {save_scaler_path}") | |
| # Normalize feature vectors | |
| normalized_features = scaler.transform(raw_features) | |
| if verbose: | |
| print("Normalized molecule features") | |
| return normalized_features, scaler | |