tox21_tabicl_classifier / preprocess.py
Queimo's picture
Convert Tox21 baseline to TabICLv2
3286d5c verified
Raw History Blame Contribute Delete
3.31 kB
# pipeline taken from https://huggingface.co/spaces/ml-jku/mhnfs/blob/main/src/data_preprocessing/create_descriptors.py
"""
This files includes a the data processing for Tox21.
As an input it takes a list of SMILES and it outputs a nested dictionary with
SMILES and target names as keys.
"""
import os
import json
import argparse
import numpy as np
from src.data import create_descriptors, get_tox21_split
from src.utils import TASKS, HF_TOKEN, write_pickle, create_dir, normalize_config
parser = argparse.ArgumentParser(
description="Data preprocessing script for the Tox21 dataset"
)
parser.add_argument(
"--config",
type=str,
default="config/config.json",
)
def main(config):
"""Preprocess the training and validation data for TabICLv2.
1. Download Tox21 train/val data from HF
2. Preprocess dataset splits
"""
ds = get_tox21_split(HF_TOKEN, cvfold=config["cvfold"])
feature_creation_kwargs = {
"radius": config["ecfp"]["radius"],
"fpsize": config["ecfp"]["fpsize"],
"min_var": config["feature_selection"]["min_var"],
"max_corr": config["feature_selection"]["max_corr"],
}
splits = ["train", "validation"]
for split in splits:
print(f"Preprocess {split} molecules")
ds_split = ds[split]
smiles = list(ds_split["smiles"])
if split == "train":
output = create_descriptors(
smiles,
return_feature_selection=True,
return_ecdfs=True,
**feature_creation_kwargs,
)
features = output.pop("features")
feature_selection = output.pop("feature_selection")
ecdfs = output.pop("ecdfs")
feature_selection_path = os.path.join(
config["data_folder"], "feat_selection.npz"
)
np.savez(
feature_selection_path,
ecfps_selec=feature_selection["ecfps_selec"],
tox_selec=feature_selection["tox_selec"],
)
print(f"Saved feature selection under {feature_selection_path}")
ecdfs_path = os.path.join(config["data_folder"], "ecdfs.pkl")
write_pickle(ecdfs_path, ecdfs)
print(f"Saved ECDFs under {ecdfs_path}")
else:
features = create_descriptors(
smiles,
ecdfs=ecdfs,
feature_selection=feature_selection,
**feature_creation_kwargs,
)["features"]
labels = []
for task in TASKS:
labels.append(ds_split[task].to_numpy())
labels = np.stack(labels, axis=1)
save_path = os.path.join(
config["data_folder"], f"tox21_{split}_cv{config['cvfold']}.npz"
)
with open(save_path, "wb") as f:
np.savez(
f,
labels=labels,
**features,
)
print(f"Saved preprocessed {split} split under {config['data_folder']}")
print("Preprocessing finished successfully")
if __name__ == "__main__":
args = parser.parse_args()
with open(args.config, "r") as f:
config = json.load(f)
config = normalize_config(config)
create_dir(config["data_folder"])
main(config)