antoniaebner's picture
upload training code
5d62705
Raw
History Blame Contribute Delete
11 kB
"""
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