Molecular Dataset Splitting Strategies: Visualization and Performance Analysis¶
Overview¶
This notebook provides comprehensive visualizations and analysis of different molecular dataset splitting strategies and their impact on machine learning model performance. We evaluate both classical ML and Graph Neural Network (GNN) models across multiple TDC datasets, comparing in-distribution (ID) vs out-of-distribution (OOD) performance.
Table of Contents¶
-
- Setup and configuration for visualization libraries
- Import ALineMol utilities and dependencies
-
- 2D UMAP embeddings of molecular datasets
- Visualization of different splitting strategies in chemical space
- Comparison of train/test distributions across splits
-
- K-nearest neighbor distance calculations
- Tanimoto similarity and TMD (Tree Mover Distance) analysis
- Statistical comparison of distance distributions across splitting methods
ID vs OOD Performance Comparison
- Box plots and heatmaps of performance differences
- Model type comparison (Classical ML vs GNN)
- Statistical significance testing
Correlation Analysis Between Splits
- Regression plots showing ID vs OOD relationships
- Split-specific and dataset-specific performance analysis
- Correlation coefficients and trend analysis
Physicochemical Property Analysis
- Radar plots of molecular properties across datasets
- Distribution analysis of drug-like properties
- Molecular weight, LogP, TPSA, and other descriptors
-
- Supplementary heatmaps and performance matrices
- Summary statistics and comparative analyses
Key Findings¶
- Comparison of 8 different splitting strategies across 8 TDC datasets
- Performance evaluation using ROC-AUC, PR-AUC, and Accuracy metrics
- Analysis of chemical space coverage and model generalization
- Insights into the relationship between molecular similarity and model performance
Import Libraries¶
In [ ]:
Copied!
import os
import sys
from pathlib import Path
import json
import numpy as np
import pandas as pd
from math import pi
import matplotlib
import matplotlib.pyplot as plt
import matplotlib.gridspec as gridspec
from matplotlib.gridspec import GridSpec
from matplotlib.lines import Line2D
from scipy.stats import pearsonr
import seaborn as sns
import yaml
from sklearn.preprocessing import StandardScaler
from typing import List
from rdkit import Chem
from rdkit.Chem import Descriptors
from alinemol.utils import compare_rankings
# Resolve the repo root by walking up to the pyproject.toml marker, so this
# notebook works no matter which directory it is launched from.
repo_path = str(next(p for p in [Path.cwd(), *Path.cwd().parents] if (p / "pyproject.toml").exists()))
CHECKOUT_PATH = repo_path
DATASET_PATH = os.path.join(repo_path, "datasets")
os.chdir(CHECKOUT_PATH)
sys.path.insert(0, CHECKOUT_PATH)
import os
import sys
from pathlib import Path
import json
import numpy as np
import pandas as pd
from math import pi
import matplotlib
import matplotlib.pyplot as plt
import matplotlib.gridspec as gridspec
from matplotlib.gridspec import GridSpec
from matplotlib.lines import Line2D
from scipy.stats import pearsonr
import seaborn as sns
import yaml
from sklearn.preprocessing import StandardScaler
from typing import List
from rdkit import Chem
from rdkit.Chem import Descriptors
from alinemol.utils import compare_rankings
# Resolve the repo root by walking up to the pyproject.toml marker, so this
# notebook works no matter which directory it is launched from.
repo_path = str(next(p for p in [Path.cwd(), *Path.cwd().parents] if (p / "pyproject.toml").exists()))
CHECKOUT_PATH = repo_path
DATASET_PATH = os.path.join(repo_path, "datasets")
os.chdir(CHECKOUT_PATH)
sys.path.insert(0, CHECKOUT_PATH)
In [ ]:
Copied!
# from alinemol.splitters import ScaffoldSplit, RandomSplit, SphereExclusionSplit, KMeansSplit, DBScanSplit, OptiSimSplit
# light_color = plt.get_cmap("plasma").colors[170]
# dark_color = plt.get_cmap("plasma").colors[5]
# dark_color = "black"
# matplotlib.use("pgf")
# Set matplotlib parameters
rcparams = {
# LaTeX setup
"pgf.texsystem": "pdflatex",
"text.usetex": True,
"pgf.rcfonts": False,
# Font settings
"font.family": "serif",
"font.serif": ["Computer Modern Roman"],
"font.size": 16,
# Figure settings
"figure.dpi": 300, # Higher DPI for better quality
"figure.figsize": [6.4, 4.8], # Default figure size
"figure.constrained_layout.use": True, # Better layout handling
# Axes settings
"axes.linewidth": 1.0,
"axes.labelsize": 14,
"axes.titlesize": 14,
# Legend settings
"legend.fontsize": 14,
"legend.frameon": True,
"legend.loc": "upper right",
# Tick settings
"xtick.major.width": 1.0,
"ytick.major.width": 1.0,
"xtick.labelsize": 12,
"ytick.labelsize": 12,
}
matplotlib.rcParams.update(rcparams)
# Seaborn settings
sns.set_style("whitegrid", rc=rcparams)
sns.set_palette("Set2")
sns.set_context("paper", font_scale=1.5)
# from alinemol.splitters import ScaffoldSplit, RandomSplit, SphereExclusionSplit, KMeansSplit, DBScanSplit, OptiSimSplit
# light_color = plt.get_cmap("plasma").colors[170]
# dark_color = plt.get_cmap("plasma").colors[5]
# dark_color = "black"
# matplotlib.use("pgf")
# Set matplotlib parameters
rcparams = {
# LaTeX setup
"pgf.texsystem": "pdflatex",
"text.usetex": True,
"pgf.rcfonts": False,
# Font settings
"font.family": "serif",
"font.serif": ["Computer Modern Roman"],
"font.size": 16,
# Figure settings
"figure.dpi": 300, # Higher DPI for better quality
"figure.figsize": [6.4, 4.8], # Default figure size
"figure.constrained_layout.use": True, # Better layout handling
# Axes settings
"axes.linewidth": 1.0,
"axes.labelsize": 14,
"axes.titlesize": 14,
# Legend settings
"legend.fontsize": 14,
"legend.frameon": True,
"legend.loc": "upper right",
# Tick settings
"xtick.major.width": 1.0,
"ytick.major.width": 1.0,
"xtick.labelsize": 12,
"ytick.labelsize": 12,
}
matplotlib.rcParams.update(rcparams)
# Seaborn settings
sns.set_style("whitegrid", rc=rcparams)
sns.set_palette("Set2")
sns.set_context("paper", font_scale=1.5)
In [ ]:
Copied!
# Load the configuration file (wich contains datasets, models, and splitting)
CFG = yaml.safe_load(open(os.path.join(DATASET_PATH, "config.yml"), "r"))
ML_MODELS: List = CFG["models"]["ML"]
GNN_MODELS: List = CFG["models"]["GNN"]["scratch"]
PRETRAINED_GNN_MODELS: List = CFG["models"]["GNN"]["pretrained"]
ALL_MODELS: List = [ML_MODELS, GNN_MODELS, PRETRAINED_GNN_MODELS]
DATASET_NAMES: List = CFG["datasets"]["TDC"]
SPLIT_TYPES: List = CFG["splitting"]
SPLIT_TYPPE_MAPPING = {
"random": "Random",
"scaffold": "Scaffold",
"scaffold_generic": "Scaffold generic",
"molecular_weight": "Molecular weight",
"molecular_weight_reverse": "Molecular weight reverse",
"molecular_logp": "Molecular logP",
"kmeans": "K-means",
"max_dissimilarity": "Max dissimilarity",
"umap": "UMAP",
"hi": "Lo-Hi",
"datasail": "DataSAIL",
}
MODEL_MAPPING = {
"randomForest": "Random Forest",
"XGB": "XGBoost",
"SVM": "SVM",
"GCN": "GCN",
"GAT": "GAT",
"MPNN": "MPNN",
"AttentiveFP": "AttentiveFP",
"Weave": "Weave",
"gin_supervised_edgepred": "GIN + Edge",
"gin_supervised_contextpred": "GIN + Context",
"gin_supervised_infomax": "GIN + InfoMax",
"gin_supervised_masking": "GIN + Masking",
"gem": "GEM",
"grover": "GROVER",
}
# read the results that are saved in the results folder. This is used for the visualization
benchmark_results = pd.read_csv(os.path.join("classification_results", "TDC", "results.csv"))
# read the hit_rate csv file
hit_rate = pd.read_csv(os.path.join("classification_results", "TDC", "hit_rate.csv"))
# benchmark_results["model_type"] = benchmark_results["model"].apply(lambda x: "Classical_ML" if x in ML_MODELS else "GNN")
benchmark_results["model_type"] = benchmark_results["model"].apply(
lambda x: "Classical ML" if x in ML_MODELS else ("GNN" if x in GNN_MODELS else "Pretrained GNN")
)
hit_rate["model_type"] = hit_rate["model"].apply(
lambda x: "Classical ML" if x in ML_MODELS else ("GNN" if x in GNN_MODELS else "Pretrained GNN")
)
metric_mapping = {"accuracy": "Accuracy", "roc_auc": "ROC-AUC", "pr_auc": "PR-AUC"}
# Load the configuration file (wich contains datasets, models, and splitting)
CFG = yaml.safe_load(open(os.path.join(DATASET_PATH, "config.yml"), "r"))
ML_MODELS: List = CFG["models"]["ML"]
GNN_MODELS: List = CFG["models"]["GNN"]["scratch"]
PRETRAINED_GNN_MODELS: List = CFG["models"]["GNN"]["pretrained"]
ALL_MODELS: List = [ML_MODELS, GNN_MODELS, PRETRAINED_GNN_MODELS]
DATASET_NAMES: List = CFG["datasets"]["TDC"]
SPLIT_TYPES: List = CFG["splitting"]
SPLIT_TYPPE_MAPPING = {
"random": "Random",
"scaffold": "Scaffold",
"scaffold_generic": "Scaffold generic",
"molecular_weight": "Molecular weight",
"molecular_weight_reverse": "Molecular weight reverse",
"molecular_logp": "Molecular logP",
"kmeans": "K-means",
"max_dissimilarity": "Max dissimilarity",
"umap": "UMAP",
"hi": "Lo-Hi",
"datasail": "DataSAIL",
}
MODEL_MAPPING = {
"randomForest": "Random Forest",
"XGB": "XGBoost",
"SVM": "SVM",
"GCN": "GCN",
"GAT": "GAT",
"MPNN": "MPNN",
"AttentiveFP": "AttentiveFP",
"Weave": "Weave",
"gin_supervised_edgepred": "GIN + Edge",
"gin_supervised_contextpred": "GIN + Context",
"gin_supervised_infomax": "GIN + InfoMax",
"gin_supervised_masking": "GIN + Masking",
"gem": "GEM",
"grover": "GROVER",
}
# read the results that are saved in the results folder. This is used for the visualization
benchmark_results = pd.read_csv(os.path.join("classification_results", "TDC", "results.csv"))
# read the hit_rate csv file
hit_rate = pd.read_csv(os.path.join("classification_results", "TDC", "hit_rate.csv"))
# benchmark_results["model_type"] = benchmark_results["model"].apply(lambda x: "Classical_ML" if x in ML_MODELS else "GNN")
benchmark_results["model_type"] = benchmark_results["model"].apply(
lambda x: "Classical ML" if x in ML_MODELS else ("GNN" if x in GNN_MODELS else "Pretrained GNN")
)
hit_rate["model_type"] = hit_rate["model"].apply(
lambda x: "Classical ML" if x in ML_MODELS else ("GNN" if x in GNN_MODELS else "Pretrained GNN")
)
metric_mapping = {"accuracy": "Accuracy", "roc_auc": "ROC-AUC", "pr_auc": "PR-AUC"}
In [ ]:
Copied!
def format_xticklabels(ax, rotation=45, ha="right", fontsize=16, is_heatmap=False):
xticklabels = [label.get_text() for label in ax.get_xticklabels()]
if not is_heatmap:
ax.set_xticks(range(len(xticklabels)))
ax.set_xticklabels(
[r"\textbf{" + label + "}" for label in xticklabels], rotation=rotation, ha=ha, fontsize=fontsize
)
def format_xticklabels(ax, rotation=45, ha="right", fontsize=16, is_heatmap=False):
xticklabels = [label.get_text() for label in ax.get_xticklabels()]
if not is_heatmap:
ax.set_xticks(range(len(xticklabels)))
ax.set_xticklabels(
[r"\textbf{" + label + "}" for label in xticklabels], rotation=rotation, ha=ha, fontsize=fontsize
)
Visualize the splits in 2D¶
In [ ]:
Copied!
# Function for visulizing the chemical space (available in utils/plot_utils.py)
import umap
import datamol as dm
from typing import List
def visualize_chemspace(
data: pd.DataFrame, split_names: List[str], mol_col: str = "smiles", size_col=None, size=10, normalize=False
) -> None:
"""
Visualize chemical space using UMAP
Args:
data (pd.DataFrame): pd.DataFrame with columns "smiles", "label", "split"
split_names (list): list of split names
mol_col (str): name of column containing SMILES
size_col: name of column containing size information
Returns:
None
Note:
This UMAP embedding is based on the Morgan fingerprint of the molecules (r=2, n=2048)
"""
features = [dm.to_fp(mol) for mol in data[mol_col]]
if normalize:
features = StandardScaler().fit_transform(features)
embedding = umap.UMAP().fit_transform(features)
print(embedding.shape)
data["UMAP_0"], data["UMAP_1"] = embedding[:, 0], embedding[:, 1]
for split_name in split_names:
plt.figure(figsize=(10, 6))
fig = sns.scatterplot(
data=data,
x="UMAP_0",
y="UMAP_1",
s=size,
style=size_col,
hue=split_name,
alpha=0.5,
hue_order=["Train", "Test"],
)
fig.set_title(f"UMAP Embedding of compounds for {split_name} split", fontsize=20)
fig.legend(loc="upper right")
plt.show()
# Function for visulizing the chemical space (available in utils/plot_utils.py)
import umap
import datamol as dm
from typing import List
def visualize_chemspace(
data: pd.DataFrame, split_names: List[str], mol_col: str = "smiles", size_col=None, size=10, normalize=False
) -> None:
"""
Visualize chemical space using UMAP
Args:
data (pd.DataFrame): pd.DataFrame with columns "smiles", "label", "split"
split_names (list): list of split names
mol_col (str): name of column containing SMILES
size_col: name of column containing size information
Returns:
None
Note:
This UMAP embedding is based on the Morgan fingerprint of the molecules (r=2, n=2048)
"""
features = [dm.to_fp(mol) for mol in data[mol_col]]
if normalize:
features = StandardScaler().fit_transform(features)
embedding = umap.UMAP().fit_transform(features)
print(embedding.shape)
data["UMAP_0"], data["UMAP_1"] = embedding[:, 0], embedding[:, 1]
for split_name in split_names:
plt.figure(figsize=(10, 6))
fig = sns.scatterplot(
data=data,
x="UMAP_0",
y="UMAP_1",
s=size,
style=size_col,
hue=split_name,
alpha=0.5,
hue_order=["Train", "Test"],
)
fig.set_title(f"UMAP Embedding of compounds for {split_name} split", fontsize=20)
fig.legend(loc="upper right")
plt.show()
In [ ]:
Copied!
index = 0
dataset_categoty = "TDC"
dataset_name = "AMES"
split_types = ["scaffold", "molecular_weight", "kmeans", "max_dissimilarity", "perimeter", "molecular_logp"]
dfs = []
for i, split_type in enumerate(split_types):
# Load the dataset
train_df = pd.read_csv(
os.path.join(DATASET_PATH, dataset_categoty, dataset_name, "split", split_type, f"train_{index}.csv")
)
print(train_df.shape)
test_df = pd.read_csv(
os.path.join(DATASET_PATH, dataset_categoty, dataset_name, "split", split_type, f"test_{index}.csv")
)
print(test_df.shape)
train_test_list = ["Train"] * len(train_df) + ["Test"] * len(test_df)
df = pd.concat([train_df, test_df])
df[split_type] = train_test_list
print(df.shape)
dfs.append(df)
if i != 0:
df = pd.merge(dfs[0], df, on=["smiles", "label"], how="outer")
dfs[0] = df
index = 0
dataset_categoty = "TDC"
dataset_name = "AMES"
split_types = ["scaffold", "molecular_weight", "kmeans", "max_dissimilarity", "perimeter", "molecular_logp"]
dfs = []
for i, split_type in enumerate(split_types):
# Load the dataset
train_df = pd.read_csv(
os.path.join(DATASET_PATH, dataset_categoty, dataset_name, "split", split_type, f"train_{index}.csv")
)
print(train_df.shape)
test_df = pd.read_csv(
os.path.join(DATASET_PATH, dataset_categoty, dataset_name, "split", split_type, f"test_{index}.csv")
)
print(test_df.shape)
train_test_list = ["Train"] * len(train_df) + ["Test"] * len(test_df)
df = pd.concat([train_df, test_df])
df[split_type] = train_test_list
print(df.shape)
dfs.append(df)
if i != 0:
df = pd.merge(dfs[0], df, on=["smiles", "label"], how="outer")
dfs[0] = df
In [ ]:
Copied!
df
df
In [ ]:
Copied!
(df["scaffold"] == "Test").sum() / len(df), (df["scaffold"] == "Train").sum() / len(df)
(df["scaffold"] == "Test").sum() / len(df), (df["scaffold"] == "Train").sum() / len(df)
In [ ]:
Copied!
# import umap.plot
# umap.plot.points(features, labels=label)
# import umap.plot
# umap.plot.points(features, labels=label)
In [ ]:
Copied!
visualize_chemspace(df, split_names=split_types, size=14, normalize=False)
visualize_chemspace(df, split_names=split_types, size=14, normalize=False)
In [ ]:
Copied!
dataset_categoty = "TDC"
dataset_name = "CYP2C19"
split_type = "scaffold"
SPLIT_PATH = os.path.join(DATASET_PATH, dataset_categoty, dataset_name, "split", split_type)
train_path = os.path.join(SPLIT_PATH, "train_0.csv")
test_path = os.path.join(SPLIT_PATH, "test_0.csv")
train_df = pd.read_csv(train_path)
test_df = pd.read_csv(test_path)
dataset_categoty = "TDC"
dataset_name = "CYP2C19"
split_type = "scaffold"
SPLIT_PATH = os.path.join(DATASET_PATH, dataset_categoty, dataset_name, "split", split_type)
train_path = os.path.join(SPLIT_PATH, "train_0.csv")
test_path = os.path.join(SPLIT_PATH, "test_0.csv")
train_df = pd.read_csv(train_path)
test_df = pd.read_csv(test_path)
In [ ]:
Copied!
train_test_list = ["train"] * len(train_df) + ["test"] * len(test_df)
split_df = pd.concat([train_df, test_df])
split_df[split_type] = train_test_list
train_test_list = ["train"] * len(train_df) + ["test"] * len(test_df)
split_df = pd.concat([train_df, test_df])
split_df[split_type] = train_test_list
In [ ]:
Copied!
split_df[split_type].value_counts()
split_df[split_type].value_counts()
In [ ]:
Copied!
split_df
split_df
In [ ]:
Copied!
In [ ]:
Copied!
import splito
import splito
In [ ]:
Copied!
DF_PATH = os.path.join(DATASET_PATH, "TDC", "CYP2C19", "CYP2C19_simplified.csv")
data = pd.read_csv(DF_PATH)
data.shape
DF_PATH = os.path.join(DATASET_PATH, "TDC", "CYP2C19", "CYP2C19_simplified.csv")
data = pd.read_csv(DF_PATH)
data.shape
In [ ]:
Copied!
# Define scaffold split
splitter = splito.ScaffoldSplit(smiles=data.smiles.tolist(), n_jobs=-1, test_size=0.1, random_state=111)
# Define scaffold split
splitter = splito.ScaffoldSplit(smiles=data.smiles.tolist(), n_jobs=-1, test_size=0.1, random_state=111)
In [ ]:
Copied!
train_idx, test_idx = next(splitter.split(X=data["smiles"].values))
assert train_idx.shape[0] > test_idx.shape[0]
train_idx, test_idx = next(splitter.split(X=data["smiles"].values))
assert train_idx.shape[0] > test_idx.shape[0]
In [ ]:
Copied!
data.loc[train_idx, "ScaffoldSplit"] = "train"
data.loc[test_idx, "ScaffoldSplit"] = "test"
data["scaffold"] = splitter.scaffolds
data.loc[train_idx, "ScaffoldSplit"] = "train"
data.loc[test_idx, "ScaffoldSplit"] = "test"
data["scaffold"] = splitter.scaffolds
In [ ]:
Copied!
data["ScaffoldSplit"].value_counts()
data["ScaffoldSplit"].value_counts()
In [ ]:
Copied!
visualize_chemspace(data, split_names=["ScaffoldSplit"], size=10)
visualize_chemspace(data, split_names=["ScaffoldSplit"], size=10)
In [ ]:
Copied!
data
data
In [ ]:
Copied!
visualize_chemspace(data, split_names=["ScaffoldSplit"], mol_col="scaffold")
visualize_chemspace(data, split_names=["ScaffoldSplit"], mol_col="scaffold")
In [ ]:
Copied!
# Define PerimeterSplit
splitter = splito.PerimeterSplit(n_jobs=-1, test_size=0.2, random_state=111)
train_idx, test_idx = next(splitter.split(X=data["smiles"].values))
assert train_idx.shape[0] > test_idx.shape[0]
data.loc[train_idx, splito.PerimeterSplit.__name__] = "train"
data.loc[test_idx, splito.PerimeterSplit.__name__] = "test"
# Define PerimeterSplit
splitter = splito.PerimeterSplit(n_jobs=-1, test_size=0.2, random_state=111)
train_idx, test_idx = next(splitter.split(X=data["smiles"].values))
assert train_idx.shape[0] > test_idx.shape[0]
data.loc[train_idx, splito.PerimeterSplit.__name__] = "train"
data.loc[test_idx, splito.PerimeterSplit.__name__] = "test"
In [ ]:
Copied!
# Define PerimeterSplit
splitter = splito.MaxDissimilaritySplit(n_jobs=-1, test_size=0.2, random_state=111)
train_idx, test_idx = next(splitter.split(X=data.smiles.values))
assert train_idx.shape[0] > test_idx.shape[0]
data.loc[train_idx, "MaxDissimilaritySplit"] = "train"
data.loc[test_idx, "MaxDissimilaritySplit"] = "test"
# Define PerimeterSplit
splitter = splito.MaxDissimilaritySplit(n_jobs=-1, test_size=0.2, random_state=111)
train_idx, test_idx = next(splitter.split(X=data.smiles.values))
assert train_idx.shape[0] > test_idx.shape[0]
data.loc[train_idx, "MaxDissimilaritySplit"] = "train"
data.loc[test_idx, "MaxDissimilaritySplit"] = "test"
In [ ]:
Copied!
visualize_chemspace(data, split_names=["PerimeterSplit", "MaxDissimilaritySplit"])
visualize_chemspace(data, split_names=["PerimeterSplit", "MaxDissimilaritySplit"])
K- Nearest Dataset Distance¶
In [ ]:
Copied!
from alinemol.utils.split_utils import (
retrieve_k_nearest_neighbors_TMD,
retrieve_k_nearest_neighbors_Tanimoto,
sklearn_stratified_random_split,
)
from tqdm import tqdm
cfg = yaml.safe_load(open(os.path.join("datasets", "config.yml"), "r"))
datasets = cfg["datasets"]["TDC"]
splitting = cfg["splitting"]
# splitting = ["scaffold"]
index = np.arange(10)
df_list = []
for split in tqdm(splitting):
for dataset in datasets:
original_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", dataset, f"{dataset}_standardize.csv"))
pairwise_distance_jaccard = np.load(os.path.join(DATASET_PATH, "TDC", dataset, "Jaccard_distance.npy"))
pairwise_distance_tmd = np.load(os.path.join(DATASET_PATH, "TDC", dataset, "TMD_distance.npy"))
for idx in index:
id_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", dataset, "split", f"{split}", f"train_{idx}.csv"))
X, y = np.array(id_df["smiles"]), np.array(id_df["label"])
X_train, X_val, X_test, y_train, y_val, y_test = sklearn_stratified_random_split(X, y, (0.72, 0.08, 0.2))
train_df = pd.DataFrame({"smiles": X_train, "label": y_train})
test_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", dataset, "split", f"{split}", f"test_{idx}.csv"))
tanimoto_vals = retrieve_k_nearest_neighbors_Tanimoto(
pairwise_distance_jaccard, original_df, train_df, test_df, k=5
)
tmd_vals = retrieve_k_nearest_neighbors_TMD(pairwise_distance_tmd, original_df, train_df, test_df, k=5)
dist_df = pd.DataFrame(
{"tanimoto": tanimoto_vals, "tmd": tmd_vals, "split": split, "dataset": dataset, "index": idx}
)
df_list.append(dist_df)
dist_df = pd.concat(df_list)
dist_df
from alinemol.utils.split_utils import (
retrieve_k_nearest_neighbors_TMD,
retrieve_k_nearest_neighbors_Tanimoto,
sklearn_stratified_random_split,
)
from tqdm import tqdm
cfg = yaml.safe_load(open(os.path.join("datasets", "config.yml"), "r"))
datasets = cfg["datasets"]["TDC"]
splitting = cfg["splitting"]
# splitting = ["scaffold"]
index = np.arange(10)
df_list = []
for split in tqdm(splitting):
for dataset in datasets:
original_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", dataset, f"{dataset}_standardize.csv"))
pairwise_distance_jaccard = np.load(os.path.join(DATASET_PATH, "TDC", dataset, "Jaccard_distance.npy"))
pairwise_distance_tmd = np.load(os.path.join(DATASET_PATH, "TDC", dataset, "TMD_distance.npy"))
for idx in index:
id_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", dataset, "split", f"{split}", f"train_{idx}.csv"))
X, y = np.array(id_df["smiles"]), np.array(id_df["label"])
X_train, X_val, X_test, y_train, y_val, y_test = sklearn_stratified_random_split(X, y, (0.72, 0.08, 0.2))
train_df = pd.DataFrame({"smiles": X_train, "label": y_train})
test_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", dataset, "split", f"{split}", f"test_{idx}.csv"))
tanimoto_vals = retrieve_k_nearest_neighbors_Tanimoto(
pairwise_distance_jaccard, original_df, train_df, test_df, k=5
)
tmd_vals = retrieve_k_nearest_neighbors_TMD(pairwise_distance_tmd, original_df, train_df, test_df, k=5)
dist_df = pd.DataFrame(
{"tanimoto": tanimoto_vals, "tmd": tmd_vals, "split": split, "dataset": dataset, "index": idx}
)
df_list.append(dist_df)
dist_df = pd.concat(df_list)
dist_df
In [ ]:
Copied!
dist_df.to_csv(os.path.join(DATASET_PATH, "TDC", "all_splitters_nearest_distances.csv"))
dist_df.to_csv(os.path.join(DATASET_PATH, "TDC", "all_splitters_nearest_distances.csv"))
In [ ]:
Copied!
dist_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", "all_splitters_nearest_distances.csv"))
dist_df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", "all_splitters_nearest_distances.csv"))
In [ ]:
Copied!
# First, group based on split and compute Median for each split for jaccard distance. Seond, group based on split and compute Median for each split for TMD distance
jaccard_df = dist_df.groupby(["split"])["tanimoto"].median().reset_index()
tmd_df = dist_df.groupby(["split"])["tmd"].median().reset_index()
# First, group based on split and compute Median for each split for jaccard distance. Seond, group based on split and compute Median for each split for TMD distance
jaccard_df = dist_df.groupby(["split"])["tanimoto"].median().reset_index()
tmd_df = dist_df.groupby(["split"])["tmd"].median().reset_index()
In [ ]:
Copied!
categories = jaccard_df["split"].tolist()
condition1 = jaccard_df["tanimoto"].tolist()
condition2 = tmd_df["tmd"].tolist()
# Compare rankings
results = compare_rankings(condition1, condition2, categories)
# Print results
print("\nRanking Comparison Results:")
print(f"Spearman Correlation: {results['spearman_correlation']:.3f} (p-value: {results['spearman_p_value']:.3f})")
print(f"Kendall's Tau: {results['kendall_tau']:.3f} (p-value: {results['kendall_p_value']:.3f})")
print(f"Footrule Distance: {results['footrule_distance']:.0f}")
print("\nDetailed Comparison:")
print(results["comparison_table"].to_string(index=False))
categories = jaccard_df["split"].tolist()
condition1 = jaccard_df["tanimoto"].tolist()
condition2 = tmd_df["tmd"].tolist()
# Compare rankings
results = compare_rankings(condition1, condition2, categories)
# Print results
print("\nRanking Comparison Results:")
print(f"Spearman Correlation: {results['spearman_correlation']:.3f} (p-value: {results['spearman_p_value']:.3f})")
print(f"Kendall's Tau: {results['kendall_tau']:.3f} (p-value: {results['kendall_p_value']:.3f})")
print(f"Footrule Distance: {results['footrule_distance']:.0f}")
print("\nDetailed Comparison:")
print(results["comparison_table"].to_string(index=False))
In [ ]:
Copied!
# Plot a Boxplot Figure (1, 2) with two horizontal axis, one for Tanimoto and one for TMD. (x-axis: split, y-axis: distance)
save = False
fig, ax = plt.subplots(1, 2, figsize=(14, 6), gridspec_kw={"wspace": 0.1})
# map split names to their full names
dist_df_plot = dist_df.copy()
dist_df_plot["split"] = dist_df_plot["split"].map(SPLIT_TYPPE_MAPPING)
sns.boxplot(
x="split",
y="tanimoto",
data=dist_df_plot,
hue="split",
ax=ax[0],
showfliers=False,
medianprops={"linewidth": 1.5},
width=0.7,
capwidths=0.25,
)
sns.boxplot(
x="split",
y="tmd",
data=dist_df_plot,
hue="split",
ax=ax[1],
showfliers=False,
medianprops={"linewidth": 1.5},
width=0.7,
capwidths=0.25,
)
# ax[0].set_title("Tanimoto Distance to 5 nearest neighbors", fontsize=20)
# ax[1].set_title("TMD Distance to 5 nearest neighbors", fontsize=20)
ax[0].set_xlabel("", fontsize=20)
ax[1].set_xlabel("", fontsize=20)
ax[0].set_ylabel(r"$\textbf{Tanimoto Distance}$", fontsize=20)
ax[1].set_ylabel(r"$\textbf{TMD}$", fontsize=20)
ax[0].grid(axis="both", linestyle="--", alpha=0.6)
ax[1].grid(axis="both", linestyle="--", alpha=0.6)
format_xticklabels(ax[0])
format_xticklabels(ax[1])
plt.setp(ax[0].get_yticklabels(), fontsize=16)
plt.setp(ax[1].get_yticklabels(), fontsize=16)
# put a and b in the top left corner of the first and second subplot (outside the plot)
ax[0].text(-0.12, 1.05, r"$\textbf{(a)}$", transform=ax[0].transAxes, fontsize=20)
ax[1].text(-0.15, 1.05, r"$\textbf{(b)}$", transform=ax[1].transAxes, fontsize=20)
# save to pdf and png
if save:
plt.savefig("assets/figures/distance_distribution.pdf", bbox_inches="tight")
plt.savefig("assets/figures/distance_distribution.png", bbox_inches="tight", dpi=600)
plt.show()
# Plot a Boxplot Figure (1, 2) with two horizontal axis, one for Tanimoto and one for TMD. (x-axis: split, y-axis: distance)
save = False
fig, ax = plt.subplots(1, 2, figsize=(14, 6), gridspec_kw={"wspace": 0.1})
# map split names to their full names
dist_df_plot = dist_df.copy()
dist_df_plot["split"] = dist_df_plot["split"].map(SPLIT_TYPPE_MAPPING)
sns.boxplot(
x="split",
y="tanimoto",
data=dist_df_plot,
hue="split",
ax=ax[0],
showfliers=False,
medianprops={"linewidth": 1.5},
width=0.7,
capwidths=0.25,
)
sns.boxplot(
x="split",
y="tmd",
data=dist_df_plot,
hue="split",
ax=ax[1],
showfliers=False,
medianprops={"linewidth": 1.5},
width=0.7,
capwidths=0.25,
)
# ax[0].set_title("Tanimoto Distance to 5 nearest neighbors", fontsize=20)
# ax[1].set_title("TMD Distance to 5 nearest neighbors", fontsize=20)
ax[0].set_xlabel("", fontsize=20)
ax[1].set_xlabel("", fontsize=20)
ax[0].set_ylabel(r"$\textbf{Tanimoto Distance}$", fontsize=20)
ax[1].set_ylabel(r"$\textbf{TMD}$", fontsize=20)
ax[0].grid(axis="both", linestyle="--", alpha=0.6)
ax[1].grid(axis="both", linestyle="--", alpha=0.6)
format_xticklabels(ax[0])
format_xticklabels(ax[1])
plt.setp(ax[0].get_yticklabels(), fontsize=16)
plt.setp(ax[1].get_yticklabels(), fontsize=16)
# put a and b in the top left corner of the first and second subplot (outside the plot)
ax[0].text(-0.12, 1.05, r"$\textbf{(a)}$", transform=ax[0].transAxes, fontsize=20)
ax[1].text(-0.15, 1.05, r"$\textbf{(b)}$", transform=ax[1].transAxes, fontsize=20)
# save to pdf and png
if save:
plt.savefig("assets/figures/distance_distribution.pdf", bbox_inches="tight")
plt.savefig("assets/figures/distance_distribution.png", bbox_inches="tight", dpi=600)
plt.show()
Plot performance of ID vs OOD¶
In [ ]:
Copied!
from alinemol.utils.plot_utils import boxplot_heatmap_performance_difference
boxplot_heatmap_performance_difference(results=benchmark_results, metric="roc_auc", save=True)
from alinemol.utils.plot_utils import boxplot_heatmap_performance_difference
boxplot_heatmap_performance_difference(results=benchmark_results, metric="roc_auc", save=True)
In [ ]:
Copied!
# First groupby by the split. Then box plot the differnce between ID_test_roc_auc and OOD_test_roc_auc for each split
# Create a single figure with two subplots
# box plots with the difference between ID and OOD test aggregated by dataset
# heatmap with the difference between ID and OOD test aggregated by dataset and split
metric = "roc_auc"
perc = False
save = False
metric_mapping = {"accuracy": "Accuracy", "roc_auc": "ROC-AUC", "pr_auc": "PR-AUC"}
diff = benchmark_results[f"ID_test_{metric}"] - benchmark_results[f"OOD_test_{metric}"]
benchmark_results["difference"] = diff
benchmark_results_plot = benchmark_results.copy()
benchmark_results_plot["split"] = benchmark_results_plot["split"].map(SPLIT_TYPPE_MAPPING)
# Create a single figure with two subplots with more horizontal space between them
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6), gridspec_kw={"wspace": 0.1})
# First subplot - Boxplot
sns.boxplot(
x="split",
y="difference",
data=benchmark_results_plot,
hue="split",
ax=ax1,
showfliers=True,
medianprops={"linewidth": 1.5},
width=0.7,
capwidths=0.25,
fliersize=1,
)
# plt.setp(ax1.get_xticklabels(), rotation=60, ha='right', fontsize=16)
format_xticklabels(ax1)
plt.setp(ax1.get_yticklabels(), fontsize=16)
ax1.set_title(r"$\textbf{Difference between ID and OOD test " + metric_mapping[metric] + "}$", fontsize=20, pad=20)
ax1.set_xlabel("", fontsize=24)
ax1.set_ylabel("$\Delta$" + r"$\textbf{" + metric_mapping[metric] + "}$", fontsize=20)
ax1.grid(axis="y", linestyle="--", alpha=0.6)
# ax1.legend().remove() # Remove redundant legend
# Second subplot - Heatmap
dataset_names = benchmark_results_plot["dataset"].unique()
split_types = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df = pd.DataFrame(index=dataset_names, columns=split_types)
# Fill the dataframe
for dataset in dataset_names:
for split in split_types:
num = benchmark_results_plot[
(benchmark_results_plot["dataset"] == dataset) & (benchmark_results_plot["split"] == split)
][f"ID_test_{metric}"].mean()
den = benchmark_results_plot[
(benchmark_results_plot["dataset"] == dataset) & (benchmark_results_plot["split"] == split)
][f"OOD_test_{metric}"].mean()
df.loc[dataset, split] = num - den if not perc else (num - den) / num * 100
df = df.astype(float)
# Plot heatmap
vmin, vmax = 0.0, 0.2
sns.heatmap(df, ax=ax2, cmap="coolwarm", annot=True, fmt=".3f", vmin=vmin, vmax=vmax, annot_kws={"size": 12})
ax2.set_xlabel("", fontsize=24)
ax2.set_ylabel(r"$\textbf{Data Sets}$", fontsize=20)
ax2.set_title(r"$\textbf{Difference between ID and OOD test " + metric_mapping[metric] + "}$", fontsize=20, pad=20)
format_xticklabels(ax2, rotation=45, ha="right", fontsize=16, is_heatmap=True)
plt.setp(ax2.get_yticklabels(), fontsize=16, ha="right")
ax2.collections[0].colorbar.ax.tick_params(labelsize=12)
# Add subplot labels
ax1.text(-0.15, 1.2, r"$\textbf{(a)}$", transform=ax1.transAxes, fontsize=20)
ax2.text(-0.23, 1.2, r"$\textbf{(b)}$", transform=ax2.transAxes, fontsize=20)
# Adjust layout with bottom margin
# plt.subplots_adjust(bottom=0.2)
# Save if needed
if save:
plt.savefig("assets/box_heatmap_id_ood_comparison_roc_auc.pdf", bbox_inches="tight")
plt.savefig("assets/box_heatmap_id_ood_comparison_roc_auc.png", bbox_inches="tight", dpi=600)
plt.show()
# First groupby by the split. Then box plot the differnce between ID_test_roc_auc and OOD_test_roc_auc for each split
# Create a single figure with two subplots
# box plots with the difference between ID and OOD test aggregated by dataset
# heatmap with the difference between ID and OOD test aggregated by dataset and split
metric = "roc_auc"
perc = False
save = False
metric_mapping = {"accuracy": "Accuracy", "roc_auc": "ROC-AUC", "pr_auc": "PR-AUC"}
diff = benchmark_results[f"ID_test_{metric}"] - benchmark_results[f"OOD_test_{metric}"]
benchmark_results["difference"] = diff
benchmark_results_plot = benchmark_results.copy()
benchmark_results_plot["split"] = benchmark_results_plot["split"].map(SPLIT_TYPPE_MAPPING)
# Create a single figure with two subplots with more horizontal space between them
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6), gridspec_kw={"wspace": 0.1})
# First subplot - Boxplot
sns.boxplot(
x="split",
y="difference",
data=benchmark_results_plot,
hue="split",
ax=ax1,
showfliers=True,
medianprops={"linewidth": 1.5},
width=0.7,
capwidths=0.25,
fliersize=1,
)
# plt.setp(ax1.get_xticklabels(), rotation=60, ha='right', fontsize=16)
format_xticklabels(ax1)
plt.setp(ax1.get_yticklabels(), fontsize=16)
ax1.set_title(r"$\textbf{Difference between ID and OOD test " + metric_mapping[metric] + "}$", fontsize=20, pad=20)
ax1.set_xlabel("", fontsize=24)
ax1.set_ylabel("$\Delta$" + r"$\textbf{" + metric_mapping[metric] + "}$", fontsize=20)
ax1.grid(axis="y", linestyle="--", alpha=0.6)
# ax1.legend().remove() # Remove redundant legend
# Second subplot - Heatmap
dataset_names = benchmark_results_plot["dataset"].unique()
split_types = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df = pd.DataFrame(index=dataset_names, columns=split_types)
# Fill the dataframe
for dataset in dataset_names:
for split in split_types:
num = benchmark_results_plot[
(benchmark_results_plot["dataset"] == dataset) & (benchmark_results_plot["split"] == split)
][f"ID_test_{metric}"].mean()
den = benchmark_results_plot[
(benchmark_results_plot["dataset"] == dataset) & (benchmark_results_plot["split"] == split)
][f"OOD_test_{metric}"].mean()
df.loc[dataset, split] = num - den if not perc else (num - den) / num * 100
df = df.astype(float)
# Plot heatmap
vmin, vmax = 0.0, 0.2
sns.heatmap(df, ax=ax2, cmap="coolwarm", annot=True, fmt=".3f", vmin=vmin, vmax=vmax, annot_kws={"size": 12})
ax2.set_xlabel("", fontsize=24)
ax2.set_ylabel(r"$\textbf{Data Sets}$", fontsize=20)
ax2.set_title(r"$\textbf{Difference between ID and OOD test " + metric_mapping[metric] + "}$", fontsize=20, pad=20)
format_xticklabels(ax2, rotation=45, ha="right", fontsize=16, is_heatmap=True)
plt.setp(ax2.get_yticklabels(), fontsize=16, ha="right")
ax2.collections[0].colorbar.ax.tick_params(labelsize=12)
# Add subplot labels
ax1.text(-0.15, 1.2, r"$\textbf{(a)}$", transform=ax1.transAxes, fontsize=20)
ax2.text(-0.23, 1.2, r"$\textbf{(b)}$", transform=ax2.transAxes, fontsize=20)
# Adjust layout with bottom margin
# plt.subplots_adjust(bottom=0.2)
# Save if needed
if save:
plt.savefig("assets/box_heatmap_id_ood_comparison_roc_auc.pdf", bbox_inches="tight")
plt.savefig("assets/box_heatmap_id_ood_comparison_roc_auc.png", bbox_inches="tight", dpi=600)
plt.show()
In [ ]:
Copied!
from alinemol.utils.plot_utils import barplot_performance_difference_model_categories
barplot_performance_difference_model_categories(results=benchmark_results, metric="roc_auc", save=True)
from alinemol.utils.plot_utils import barplot_performance_difference_model_categories
barplot_performance_difference_model_categories(results=benchmark_results, metric="roc_auc", save=True)
In [ ]:
Copied!
# Create grouped bar plot for ML vs GNN model differences
save = False
fig = plt.figure(figsize=(16, 12))
# gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, wspace=0.5, figure=fig) # Increased hspace for more gap
gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, figure=fig) # Increased hspace for more gap
# First row: one wide subplot
gs1 = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=gs[0])
ax1 = plt.subplot(gs1[0, 0])
# Second and third rows: 4x2 grid (8 subplots)
gs2 = gridspec.GridSpecFromSubplotSpec(2, 4, subplot_spec=gs[1])
axes = [plt.subplot(gs2[i, j]) for i in range(2) for j in range(4)]
df = benchmark_results.groupby(["split", "model_type", "dataset"])["difference"].mean().reset_index()
df["split"] = df["split"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["split"] = pd.Categorical(df["split"], split_order)
# Main plot
g = sns.barplot(x="split", y="difference", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_xlabel("") # No x-label on main plot
ax1.set_ylabel(r"\textbf{$\Delta$ " + metric_mapping[metric] + "}", fontsize=20)
ax1.set_title(r"\textbf{Difference between ID and OOD test " + metric_mapping[metric] + "}", fontsize=24, pad=20)
# ticks formatting
format_xticklabels(ax1, rotation=30, ha="right", fontsize=16)
plt.setp(ax1.get_yticklabels(), fontsize=16)
# Grid and legend
ax1.grid(axis="y", linestyle="--", alpha=0.6)
legend = ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=14, title=r"\textbf{Model Type}")
for text in legend.get_texts():
text.set_fontsize(14)
text.set_text(r"\textbf{" + text.get_text() + "}")
# Subplots
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="split", y="difference", hue="model_type", data=df_dataset, ax=axes[i])
axes[i].set_title(r"\textbf{" + dataset + "}", fontsize=18, pad=10)
axes[i].set_xlabel("")
if i % 4 == 0:
axes[i].set_ylabel(r"\textbf{$\Delta$ " + metric_mapping[metric] + "}", fontsize=14)
else:
axes[i].set_ylabel("")
# x-tick formatting
if i >= 4:
format_xticklabels(axes[i])
else:
axes[i].set_xticklabels([])
# y-tick formatting
# yticklabels = [label.get_text() for label in axes[i].get_yticklabels()]
# axes[i].set_yticklabels([r'\textbf{' + label + '}' for label in yticklabels], fontsize=10)
axes[i].tick_params(axis="both", labelsize=14)
axes[i].grid(axis="y", linestyle="--", alpha=0.5)
g.legend().remove()
axes[i].set_ylim(-0.10, 0.25)
axes[i].set_yticks(np.arange(-0.10, 0.26, 0.05))
# Subplot labels
ax1.text(-0.08, 1.15, r"\textbf{(a)}", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.27, 1.35, r"\textbf{(b)}", transform=axes[0].transAxes, fontsize=20)
# plt.subplots_adjust(left=0.07, right=0.95, top=0.93, bottom=0.08, hspace=1, wspace=0.45) # Increased hspace
if save:
plt.savefig("assets/figures/grouped_barplot_ml_gnn_difference.pdf", bbox_inches="tight")
plt.savefig("assets/figures/grouped_barplot_ml_gnn_difference.png", bbox_inches="tight", dpi=600) # High-res
plt.show()
# Create grouped bar plot for ML vs GNN model differences
save = False
fig = plt.figure(figsize=(16, 12))
# gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, wspace=0.5, figure=fig) # Increased hspace for more gap
gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, figure=fig) # Increased hspace for more gap
# First row: one wide subplot
gs1 = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=gs[0])
ax1 = plt.subplot(gs1[0, 0])
# Second and third rows: 4x2 grid (8 subplots)
gs2 = gridspec.GridSpecFromSubplotSpec(2, 4, subplot_spec=gs[1])
axes = [plt.subplot(gs2[i, j]) for i in range(2) for j in range(4)]
df = benchmark_results.groupby(["split", "model_type", "dataset"])["difference"].mean().reset_index()
df["split"] = df["split"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["split"] = pd.Categorical(df["split"], split_order)
# Main plot
g = sns.barplot(x="split", y="difference", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_xlabel("") # No x-label on main plot
ax1.set_ylabel(r"\textbf{$\Delta$ " + metric_mapping[metric] + "}", fontsize=20)
ax1.set_title(r"\textbf{Difference between ID and OOD test " + metric_mapping[metric] + "}", fontsize=24, pad=20)
# ticks formatting
format_xticklabels(ax1, rotation=30, ha="right", fontsize=16)
plt.setp(ax1.get_yticklabels(), fontsize=16)
# Grid and legend
ax1.grid(axis="y", linestyle="--", alpha=0.6)
legend = ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=14, title=r"\textbf{Model Type}")
for text in legend.get_texts():
text.set_fontsize(14)
text.set_text(r"\textbf{" + text.get_text() + "}")
# Subplots
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="split", y="difference", hue="model_type", data=df_dataset, ax=axes[i])
axes[i].set_title(r"\textbf{" + dataset + "}", fontsize=18, pad=10)
axes[i].set_xlabel("")
if i % 4 == 0:
axes[i].set_ylabel(r"\textbf{$\Delta$ " + metric_mapping[metric] + "}", fontsize=14)
else:
axes[i].set_ylabel("")
# x-tick formatting
if i >= 4:
format_xticklabels(axes[i])
else:
axes[i].set_xticklabels([])
# y-tick formatting
# yticklabels = [label.get_text() for label in axes[i].get_yticklabels()]
# axes[i].set_yticklabels([r'\textbf{' + label + '}' for label in yticklabels], fontsize=10)
axes[i].tick_params(axis="both", labelsize=14)
axes[i].grid(axis="y", linestyle="--", alpha=0.5)
g.legend().remove()
axes[i].set_ylim(-0.10, 0.25)
axes[i].set_yticks(np.arange(-0.10, 0.26, 0.05))
# Subplot labels
ax1.text(-0.08, 1.15, r"\textbf{(a)}", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.27, 1.35, r"\textbf{(b)}", transform=axes[0].transAxes, fontsize=20)
# plt.subplots_adjust(left=0.07, right=0.95, top=0.93, bottom=0.08, hspace=1, wspace=0.45) # Increased hspace
if save:
plt.savefig("assets/figures/grouped_barplot_ml_gnn_difference.pdf", bbox_inches="tight")
plt.savefig("assets/figures/grouped_barplot_ml_gnn_difference.png", bbox_inches="tight", dpi=600) # High-res
plt.show()
In [ ]:
Copied!
# Create grouped bar plot for ML vs GNN model differences
save = False
benchmark_results["model_type"] = benchmark_results["model"].apply(
lambda x: "Classical_ML" if x in ML_MODELS else "GNN"
)
fig = plt.figure(figsize=(12, 10))
gs = gridspec.GridSpec(3, 4, height_ratios=[2, 1, 1]) # 3 rows, 4 columns
# First row: One wide subplot
ax1 = plt.subplot(gs[0, :]) # Spans all columns
# Second and third rows: 4x2 grid (8 subplots)
axes = [plt.subplot(gs[i, j]) for i in range(1, 3) for j in range(4)]
df = benchmark_results.groupby(["split", "model_type", "dataset"])["difference"].mean().reset_index()
df["split"] = df["split"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["split"] = pd.Categorical(df["split"], split_order)
# For ax1, plot the difference between ML and GNN models
g = sns.barplot(x="split", y="difference", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_title(f"Difference between ID and OOD test {metric_mapping[metric]}", fontsize=20)
ax1.set_xlabel("", fontsize=24)
ax1.set_ylabel(f" $\Delta$ {metric_mapping[metric]}", fontsize=16)
ax1.tick_params(axis="x", rotation=10, labelsize=10)
ax1.grid(axis="y", linestyle="--", alpha=0.6)
# Move the legend outside the plot
ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=12)
# for each dataset, create a small figure with the same plot
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="split", y="difference", hue="model_type", data=df_dataset, ax=axes[i])
# Customize the plot
# axes[i].set_title(f"Difference between ID and OOD test {metric_mapping[metric]}", fontsize=20)
# put x axis tixks just on the last rows and ignore other rows
if i < 4:
axes[i].set_xticklabels([])
else:
axes[i].tick_params(
axis="x",
rotation=90,
labelsize=12,
)
axes[i].set_title(f"{dataset}", fontsize=16)
axes[i].set_xlabel("", fontsize=18)
axes[i].set_ylabel(f"$\Delta$ {metric_mapping[metric]}", fontsize=12)
axes[i].grid(axis="both", linestyle="--", alpha=0.6)
# Adjust legend
g.legend().remove() # Remove redundant legend
# limit the range of yaxis betwen -0.1 to 0.15 with 0.05 interval
axes[i].set_ylim(-0.15, 0.15)
axes[i].set_yticks(np.arange(-0.15, 0.16, 0.05))
# remove y label for all subplots except the first one
if i % 4 != 0:
axes[i].set_ylabel("")
# add (a) to the first row, and (b) for other rows
ax1.text(-0.05, 1.3, r"$\textbf{(a)}$", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.05, 1.3, r"$\textbf{(b)}$", transform=axes[0].transAxes, fontsize=20)
# Adjust layout to prevent label cutoff
plt.tight_layout()
if save:
plt.savefig("assets/grouped_barplot_ml_gnn_difference.pdf", bbox_inches="tight")
plt.savefig("assets/grouped_barplot_ml_gnn_difference.png", bbox_inches="tight", dpi=600)
plt.show()
# Create grouped bar plot for ML vs GNN model differences
save = False
benchmark_results["model_type"] = benchmark_results["model"].apply(
lambda x: "Classical_ML" if x in ML_MODELS else "GNN"
)
fig = plt.figure(figsize=(12, 10))
gs = gridspec.GridSpec(3, 4, height_ratios=[2, 1, 1]) # 3 rows, 4 columns
# First row: One wide subplot
ax1 = plt.subplot(gs[0, :]) # Spans all columns
# Second and third rows: 4x2 grid (8 subplots)
axes = [plt.subplot(gs[i, j]) for i in range(1, 3) for j in range(4)]
df = benchmark_results.groupby(["split", "model_type", "dataset"])["difference"].mean().reset_index()
df["split"] = df["split"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["split"] = pd.Categorical(df["split"], split_order)
# For ax1, plot the difference between ML and GNN models
g = sns.barplot(x="split", y="difference", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_title(f"Difference between ID and OOD test {metric_mapping[metric]}", fontsize=20)
ax1.set_xlabel("", fontsize=24)
ax1.set_ylabel(f" $\Delta$ {metric_mapping[metric]}", fontsize=16)
ax1.tick_params(axis="x", rotation=10, labelsize=10)
ax1.grid(axis="y", linestyle="--", alpha=0.6)
# Move the legend outside the plot
ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=12)
# for each dataset, create a small figure with the same plot
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="split", y="difference", hue="model_type", data=df_dataset, ax=axes[i])
# Customize the plot
# axes[i].set_title(f"Difference between ID and OOD test {metric_mapping[metric]}", fontsize=20)
# put x axis tixks just on the last rows and ignore other rows
if i < 4:
axes[i].set_xticklabels([])
else:
axes[i].tick_params(
axis="x",
rotation=90,
labelsize=12,
)
axes[i].set_title(f"{dataset}", fontsize=16)
axes[i].set_xlabel("", fontsize=18)
axes[i].set_ylabel(f"$\Delta$ {metric_mapping[metric]}", fontsize=12)
axes[i].grid(axis="both", linestyle="--", alpha=0.6)
# Adjust legend
g.legend().remove() # Remove redundant legend
# limit the range of yaxis betwen -0.1 to 0.15 with 0.05 interval
axes[i].set_ylim(-0.15, 0.15)
axes[i].set_yticks(np.arange(-0.15, 0.16, 0.05))
# remove y label for all subplots except the first one
if i % 4 != 0:
axes[i].set_ylabel("")
# add (a) to the first row, and (b) for other rows
ax1.text(-0.05, 1.3, r"$\textbf{(a)}$", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.05, 1.3, r"$\textbf{(b)}$", transform=axes[0].transAxes, fontsize=20)
# Adjust layout to prevent label cutoff
plt.tight_layout()
if save:
plt.savefig("assets/grouped_barplot_ml_gnn_difference.pdf", bbox_inches="tight")
plt.savefig("assets/grouped_barplot_ml_gnn_difference.png", bbox_inches="tight", dpi=600)
plt.show()
In [ ]:
Copied!
for split in SPLIT_TYPES:
print(f"\nSplit: {split}")
df_subset = benchmark_results[benchmark_results["split"] == split]
# run statistical test on difference columns based on model types
from scipy.stats import ttest_ind
ml_diff = df_subset[df_subset["model_type"] == "Classical_ML"]["difference"]
gnn_diff = df_subset[df_subset["model_type"] == "GNN"]["difference"]
t_stat, p_value = ttest_ind(ml_diff, gnn_diff)
print(f"t-statistic: {t_stat:.5f}, p-value: {p_value:.5f}")
# Interpret results
alpha = 0.01 # conventional significance level
print("\nInterpretation:")
if p_value < alpha:
print(
f"Since p-value ({p_value:.4f}) is less than {alpha}, there is a significant difference between the categories"
)
else:
print(
f"Since p-value ({p_value:.4f}) is greater than {alpha}, there is no significant difference between the categories"
)
for split in SPLIT_TYPES:
print(f"\nSplit: {split}")
df_subset = benchmark_results[benchmark_results["split"] == split]
# run statistical test on difference columns based on model types
from scipy.stats import ttest_ind
ml_diff = df_subset[df_subset["model_type"] == "Classical_ML"]["difference"]
gnn_diff = df_subset[df_subset["model_type"] == "GNN"]["difference"]
t_stat, p_value = ttest_ind(ml_diff, gnn_diff)
print(f"t-statistic: {t_stat:.5f}, p-value: {p_value:.5f}")
# Interpret results
alpha = 0.01 # conventional significance level
print("\nInterpretation:")
if p_value < alpha:
print(
f"Since p-value ({p_value:.4f}) is less than {alpha}, there is a significant difference between the categories"
)
else:
print(
f"Since p-value ({p_value:.4f}) is greater than {alpha}, there is no significant difference between the categories"
)
In [ ]:
Copied!
from alinemol.utils.plot_utils import heatmap_plot_all_dataset
heatmap_plot_all_dataset(results=benchmark_results, metric="roc_auc", save=True)
from alinemol.utils.plot_utils import heatmap_plot_all_dataset
heatmap_plot_all_dataset(results=benchmark_results, metric="roc_auc", save=True)
Comparing Two or more splits together¶
Regression plot between ID and OOD performance of model categories¶
In [ ]:
Copied!
from alinemol.utils.plot_utils import regplot_with_model_categories
regplot_with_model_categories(results=benchmark_results, metric="roc_auc", save=False)
from alinemol.utils.plot_utils import regplot_with_model_categories
regplot_with_model_categories(results=benchmark_results, metric="roc_auc", save=False)
In [ ]:
Copied!
from alinemol.utils.plot_utils import regplot_with_model_categories_fixed_split
regplot_with_model_categories_fixed_split(
results=benchmark_results, split="max_dissimilarity", metric="roc_auc", save=False
)
from alinemol.utils.plot_utils import regplot_with_model_categories_fixed_split
regplot_with_model_categories_fixed_split(
results=benchmark_results, split="max_dissimilarity", metric="roc_auc", save=False
)
In [ ]:
Copied!
save = False
metric = "roc_auc"
def make_regplot(x, y, ax, color="#66c2a5", linewidth=2):
"""
Make a regression plot
"""
sns.regplot(
x=x,
y=y,
ax=ax,
scatter_kws={"alpha": 0.5, "s": 40},
line_kws={"color": "red", "linewidth": linewidth},
ci=95,
color=color,
)
corr = pearsonr(x, y)[0]
props = dict(boxstyle="round", facecolor="white", alpha=0.8, edgecolor="gray")
ax.text(0.05, 0.95, f"r = {corr:.2f}", transform=ax.transAxes, fontsize=14, verticalalignment="top", bbox=props)
def format_axis(ax, linewidth=2, ticklabelsize=14):
"""
Format the axis of the plot
"""
lims = [0.5, 0.95]
ax.set_xlim(lims)
ax.set_ylim(lims)
ax.plot(lims, lims, "--", alpha=1, color="gray", linewidth=linewidth)
ax.set_xticks(np.arange(0.5, 1.0, 0.1))
ax.set_yticks(np.arange(0.5, 1.0, 0.1))
ax.tick_params(axis="both", which="major", labelsize=ticklabelsize)
ax.set_xticklabels(ax.get_xticklabels(), fontsize=ticklabelsize)
ax.set_yticklabels(ax.get_yticklabels(), fontsize=ticklabelsize)
df = benchmark_results.copy()
df["difference"] = df[f"ID_test_{metric}"] - df[f"OOD_test_{metric}"]
fig = plt.figure(figsize=(16, 12))
gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, figure=fig) # Increased hspace for more gap
# First row: one wide subplot
gs1 = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=gs[0])
ax1 = plt.subplot(gs1[0, 0])
# Second and third rows: 4x2 grid (8 subplots)
gs2 = gridspec.GridSpecFromSubplotSpec(2, 4, subplot_spec=gs[1])
axes = [plt.subplot(gs2[i, j]) for i in range(2) for j in range(4)]
results_plot = results.copy()
results_plot["difference"] = results_plot[f"ID_test_{metric}"] - results_plot[f"OOD_test_{metric}"]
df = results_plot.groupby(["split", "model_type", "dataset"])["difference"].mean().reset_index()
df["split"] = df["split"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["split"] = pd.Categorical(df["split"], split_order)
# Main plot
g = sns.barplot(x="split", y="difference", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_xlabel("") # No x-label on main plot
ax1.set_ylabel(r"\textbf{$\Delta$ " + METRIC_MAPPING[metric] + "}", fontsize=20)
ax1.set_title(r"\textbf{Difference between ID and OOD test " + METRIC_MAPPING[metric] + "}", fontsize=24, pad=20)
# ticks formatting
format_xticklabels(ax1, rotation=30, ha="right", fontsize=16)
plt.setp(ax1.get_yticklabels(), fontsize=16)
# Grid and legend
ax1.grid(axis="y", linestyle="--", alpha=0.6)
legend = ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=14, title=r"\textbf{Model Type}")
for text in legend.get_texts():
text.set_fontsize(14)
text.set_text(r"\textbf{" + text.get_text() + "}")
# Subplots
dataset_names = results_plot["dataset"].unique()
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="split", y="difference", hue="model_type", data=df_dataset, ax=axes[i])
axes[i].set_title(r"\textbf{" + dataset + "}", fontsize=18, pad=10)
axes[i].set_xlabel("")
if i % 4 == 0:
axes[i].set_ylabel(r"\textbf{$\Delta$ " + METRIC_MAPPING[metric] + "}", fontsize=14)
else:
axes[i].set_ylabel("")
if i >= 4:
format_xticklabels(axes[i])
else:
axes[i].set_xticklabels([])
axes[i].tick_params(axis="both", labelsize=14)
axes[i].grid(axis="y", linestyle="--", alpha=0.5)
g.legend().remove()
axes[i].set_ylim(-0.10, 0.20)
axes[i].set_yticks(np.arange(-0.10, 0.21, 0.05))
# Subplot labels
ax1.text(-0.08, 1.15, r"\textbf{(a)}", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.27, 1.35, r"\textbf{(b)}", transform=axes[0].transAxes, fontsize=20)
if save:
plt.savefig(
os.path.join(REPO_PATH, "assets", "figures", "grouped_barplot_ml_gnn_difference.pdf"), bbox_inches="tight"
)
plt.savefig(
os.path.join(REPO_PATH, "assets", "figures", "grouped_barplot_ml_gnn_difference.png"),
bbox_inches="tight",
dpi=600,
)
plt.show()
save = False
metric = "roc_auc"
def make_regplot(x, y, ax, color="#66c2a5", linewidth=2):
"""
Make a regression plot
"""
sns.regplot(
x=x,
y=y,
ax=ax,
scatter_kws={"alpha": 0.5, "s": 40},
line_kws={"color": "red", "linewidth": linewidth},
ci=95,
color=color,
)
corr = pearsonr(x, y)[0]
props = dict(boxstyle="round", facecolor="white", alpha=0.8, edgecolor="gray")
ax.text(0.05, 0.95, f"r = {corr:.2f}", transform=ax.transAxes, fontsize=14, verticalalignment="top", bbox=props)
def format_axis(ax, linewidth=2, ticklabelsize=14):
"""
Format the axis of the plot
"""
lims = [0.5, 0.95]
ax.set_xlim(lims)
ax.set_ylim(lims)
ax.plot(lims, lims, "--", alpha=1, color="gray", linewidth=linewidth)
ax.set_xticks(np.arange(0.5, 1.0, 0.1))
ax.set_yticks(np.arange(0.5, 1.0, 0.1))
ax.tick_params(axis="both", which="major", labelsize=ticklabelsize)
ax.set_xticklabels(ax.get_xticklabels(), fontsize=ticklabelsize)
ax.set_yticklabels(ax.get_yticklabels(), fontsize=ticklabelsize)
df = benchmark_results.copy()
df["difference"] = df[f"ID_test_{metric}"] - df[f"OOD_test_{metric}"]
fig = plt.figure(figsize=(16, 12))
gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, figure=fig) # Increased hspace for more gap
# First row: one wide subplot
gs1 = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=gs[0])
ax1 = plt.subplot(gs1[0, 0])
# Second and third rows: 4x2 grid (8 subplots)
gs2 = gridspec.GridSpecFromSubplotSpec(2, 4, subplot_spec=gs[1])
axes = [plt.subplot(gs2[i, j]) for i in range(2) for j in range(4)]
results_plot = results.copy()
results_plot["difference"] = results_plot[f"ID_test_{metric}"] - results_plot[f"OOD_test_{metric}"]
df = results_plot.groupby(["split", "model_type", "dataset"])["difference"].mean().reset_index()
df["split"] = df["split"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["split"] = pd.Categorical(df["split"], split_order)
# Main plot
g = sns.barplot(x="split", y="difference", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_xlabel("") # No x-label on main plot
ax1.set_ylabel(r"\textbf{$\Delta$ " + METRIC_MAPPING[metric] + "}", fontsize=20)
ax1.set_title(r"\textbf{Difference between ID and OOD test " + METRIC_MAPPING[metric] + "}", fontsize=24, pad=20)
# ticks formatting
format_xticklabels(ax1, rotation=30, ha="right", fontsize=16)
plt.setp(ax1.get_yticklabels(), fontsize=16)
# Grid and legend
ax1.grid(axis="y", linestyle="--", alpha=0.6)
legend = ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=14, title=r"\textbf{Model Type}")
for text in legend.get_texts():
text.set_fontsize(14)
text.set_text(r"\textbf{" + text.get_text() + "}")
# Subplots
dataset_names = results_plot["dataset"].unique()
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="split", y="difference", hue="model_type", data=df_dataset, ax=axes[i])
axes[i].set_title(r"\textbf{" + dataset + "}", fontsize=18, pad=10)
axes[i].set_xlabel("")
if i % 4 == 0:
axes[i].set_ylabel(r"\textbf{$\Delta$ " + METRIC_MAPPING[metric] + "}", fontsize=14)
else:
axes[i].set_ylabel("")
if i >= 4:
format_xticklabels(axes[i])
else:
axes[i].set_xticklabels([])
axes[i].tick_params(axis="both", labelsize=14)
axes[i].grid(axis="y", linestyle="--", alpha=0.5)
g.legend().remove()
axes[i].set_ylim(-0.10, 0.20)
axes[i].set_yticks(np.arange(-0.10, 0.21, 0.05))
# Subplot labels
ax1.text(-0.08, 1.15, r"\textbf{(a)}", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.27, 1.35, r"\textbf{(b)}", transform=axes[0].transAxes, fontsize=20)
if save:
plt.savefig(
os.path.join(REPO_PATH, "assets", "figures", "grouped_barplot_ml_gnn_difference.pdf"), bbox_inches="tight"
)
plt.savefig(
os.path.join(REPO_PATH, "assets", "figures", "grouped_barplot_ml_gnn_difference.png"),
bbox_inches="tight",
dpi=600,
)
plt.show()
In [ ]:
Copied!
save = False
def make_regplot(x, y, ax, color="#2E86C1"):
sns.regplot(
x=x,
y=y,
ax=ax,
scatter_kws={"alpha": 0.5, "s": 20, "color": color},
line_kws={"color": "red", "linewidth": 2},
ci=95,
)
corr = pearsonr(x, y)[0]
ax.text(0.05, 0.95, f"r = {corr:.2f}", transform=ax.transAxes, fontsize=14)
def format_axis(ax):
lims = [np.min([ax.get_xlim(), ax.get_ylim()]), np.max([ax.get_xlim(), ax.get_ylim()])]
ax.set_xlim(lims)
ax.set_ylim(lims)
ax.plot(lims, lims, "k--", alpha=0.5, zorder=0)
# Create figure
fig = plt.figure(figsize=(16, 14))
gs = GridSpec(5, 4, height_ratios=[2, 1, 1, 1, 1], hspace=0.4, wspace=0.3)
# Add super title
fig.suptitle(f"ID vs OOD Performance Comparison ({metric_mapping[metric]})", fontsize=24, y=0.95)
# First row: One wide subplot for each model type
ax1 = plt.subplot(gs[0, 0:2])
ax2 = plt.subplot(gs[0, 2:4])
# Create subplots for splits
ml_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2)]
gnn_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2, 4)]
# Filter results
ML_result = benchmark_results[benchmark_results["model_type"] == "Classical_ML"]
GNN_result = benchmark_results[benchmark_results["model_type"] == "GNN"]
# Plot main comparisons
make_regplot(ML_result[f"ID_test_{metric}"], ML_result[f"OOD_test_{metric}"], ax1)
make_regplot(GNN_result[f"ID_test_{metric}"], GNN_result[f"OOD_test_{metric}"], ax2)
# Format main plots
for ax, title in [(ax1, "Classical ML Models"), (ax2, "GNN Models")]:
ax.set_title(title, fontsize=20, pad=20)
ax.set_xlabel(f"In-Distribution {metric_mapping[metric]}", fontsize=18)
ax.set_ylabel(f"Out-of-Distribution {metric_mapping[metric]}", fontsize=18)
ax.tick_params(axis="both", labelsize=16)
format_axis(ax)
# Plot and format split comparisons
for i, split in enumerate(SPLIT_TYPES):
ML_split = ML_result[ML_result["split"] == split]
GNN_split = GNN_result[GNN_result["split"] == split]
make_regplot(ML_split[f"ID_test_{metric}"], ML_split[f"OOD_test_{metric}"], ml_axes[i])
make_regplot(GNN_split[f"ID_test_{metric}"], GNN_split[f"OOD_test_{metric}"], gnn_axes[i])
# Format split plots
for ax in [ml_axes[i], gnn_axes[i]]:
ax.set_title(f"{split}", fontsize=20)
ax.tick_params(axis="both", labelsize=16)
format_axis(ax)
# Only show y-label for leftmost plots
if ax == ml_axes[i]:
ax.set_ylabel(f"OOD {metric_mapping[metric]}", fontsize=14)
# Only show x-label for bottom plots
if i >= len(SPLIT_TYPES) - 4:
ax.set_xlabel(f"ID {metric_mapping[metric]}", fontsize=14)
# Add legend
handles = [
Line2D([0], [0], marker="o", color="w", markerfacecolor="#2E86C1", markersize=10, label=split)
for split in SPLIT_TYPES
]
fig.legend(handles=handles, loc="center right", bbox_to_anchor=(0.98, 0.5), fontsize=14)
plt.tight_layout()
# Save plots
if save:
plt.savefig(
"assets/regplot_with_categories.pdf",
bbox_inches="tight",
dpi=300,
metadata={"Creator": "Your Name", "Title": "ID vs OOD Performance Comparison"},
)
plt.savefig("assets/regplot_with_categories.png", bbox_inches="tight", dpi=300)
plt.show()
save = False
def make_regplot(x, y, ax, color="#2E86C1"):
sns.regplot(
x=x,
y=y,
ax=ax,
scatter_kws={"alpha": 0.5, "s": 20, "color": color},
line_kws={"color": "red", "linewidth": 2},
ci=95,
)
corr = pearsonr(x, y)[0]
ax.text(0.05, 0.95, f"r = {corr:.2f}", transform=ax.transAxes, fontsize=14)
def format_axis(ax):
lims = [np.min([ax.get_xlim(), ax.get_ylim()]), np.max([ax.get_xlim(), ax.get_ylim()])]
ax.set_xlim(lims)
ax.set_ylim(lims)
ax.plot(lims, lims, "k--", alpha=0.5, zorder=0)
# Create figure
fig = plt.figure(figsize=(16, 14))
gs = GridSpec(5, 4, height_ratios=[2, 1, 1, 1, 1], hspace=0.4, wspace=0.3)
# Add super title
fig.suptitle(f"ID vs OOD Performance Comparison ({metric_mapping[metric]})", fontsize=24, y=0.95)
# First row: One wide subplot for each model type
ax1 = plt.subplot(gs[0, 0:2])
ax2 = plt.subplot(gs[0, 2:4])
# Create subplots for splits
ml_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2)]
gnn_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2, 4)]
# Filter results
ML_result = benchmark_results[benchmark_results["model_type"] == "Classical_ML"]
GNN_result = benchmark_results[benchmark_results["model_type"] == "GNN"]
# Plot main comparisons
make_regplot(ML_result[f"ID_test_{metric}"], ML_result[f"OOD_test_{metric}"], ax1)
make_regplot(GNN_result[f"ID_test_{metric}"], GNN_result[f"OOD_test_{metric}"], ax2)
# Format main plots
for ax, title in [(ax1, "Classical ML Models"), (ax2, "GNN Models")]:
ax.set_title(title, fontsize=20, pad=20)
ax.set_xlabel(f"In-Distribution {metric_mapping[metric]}", fontsize=18)
ax.set_ylabel(f"Out-of-Distribution {metric_mapping[metric]}", fontsize=18)
ax.tick_params(axis="both", labelsize=16)
format_axis(ax)
# Plot and format split comparisons
for i, split in enumerate(SPLIT_TYPES):
ML_split = ML_result[ML_result["split"] == split]
GNN_split = GNN_result[GNN_result["split"] == split]
make_regplot(ML_split[f"ID_test_{metric}"], ML_split[f"OOD_test_{metric}"], ml_axes[i])
make_regplot(GNN_split[f"ID_test_{metric}"], GNN_split[f"OOD_test_{metric}"], gnn_axes[i])
# Format split plots
for ax in [ml_axes[i], gnn_axes[i]]:
ax.set_title(f"{split}", fontsize=20)
ax.tick_params(axis="both", labelsize=16)
format_axis(ax)
# Only show y-label for leftmost plots
if ax == ml_axes[i]:
ax.set_ylabel(f"OOD {metric_mapping[metric]}", fontsize=14)
# Only show x-label for bottom plots
if i >= len(SPLIT_TYPES) - 4:
ax.set_xlabel(f"ID {metric_mapping[metric]}", fontsize=14)
# Add legend
handles = [
Line2D([0], [0], marker="o", color="w", markerfacecolor="#2E86C1", markersize=10, label=split)
for split in SPLIT_TYPES
]
fig.legend(handles=handles, loc="center right", bbox_to_anchor=(0.98, 0.5), fontsize=14)
plt.tight_layout()
# Save plots
if save:
plt.savefig(
"assets/regplot_with_categories.pdf",
bbox_inches="tight",
dpi=300,
metadata={"Creator": "Your Name", "Title": "ID vs OOD Performance Comparison"},
)
plt.savefig("assets/regplot_with_categories.png", bbox_inches="tight", dpi=300)
plt.show()
Regression plot between ID and OOD performance of model categories, FIXED SPLITTER¶
In [ ]:
Copied!
from alinemol.utils.plot_utils import regplot_with_model_categories_fixed_split
regplot_with_model_categories_fixed_split(
results=benchmark_results, split="max_dissimilarity", metric="roc_auc", save=True
)
from alinemol.utils.plot_utils import regplot_with_model_categories_fixed_split
regplot_with_model_categories_fixed_split(
results=benchmark_results, split="max_dissimilarity", metric="roc_auc", save=True
)
In [ ]:
Copied!
# Choose oine specific splitting strategies. Then plot for all the datasets separately the relationship between ID and OOD test roc_auc
split = "max_dissimilarity"
save = False
metric = "roc_auc"
def make_regplot(x, y, ax, color="#66c2a5"):
sns.regplot(
x=x,
y=y,
ax=ax,
scatter_kws={"alpha": 0.5, "s": 40, "color": color},
line_kws={"color": "red", "linewidth": 2},
ci=95,
)
corr = pearsonr(x, y)[0]
props = dict(boxstyle="round", facecolor="white", alpha=0.8, edgecolor="gray")
ax.text(0.05, 0.95, f"r = {corr:.2f}", transform=ax.transAxes, fontsize=14, verticalalignment="top", bbox=props)
def format_axis(ax):
# lims = [
# np.min([ax.get_xlim(), ax.get_ylim()]),
# np.max([ax.get_xlim(), ax.get_ylim()])
# ]
lims = [0.5, 0.95]
ax.set_xlim(lims)
ax.set_ylim(lims)
ax.plot(lims, lims, "b--", alpha=0.5, zorder=0)
ax.set_xticks(np.arange(0.5, 1.0, 0.1))
ax.set_yticks(np.arange(0.5, 1.0, 0.1))
fig = plt.figure(figsize=(16, 12))
gs = gridspec.GridSpec(5, 4, height_ratios=[2, 1, 1, 1, 1], hspace=1.0, wspace=0.3) # 5 rows, 4 columns
# Add super title
fig.suptitle(f"ID vs OOD Performance Comparison ({metric_mapping[metric]})", fontsize=24, y=0.95)
# First row: One wide subplot
ax1 = plt.subplot(gs[0, 0:2]) # Spans two columns
ax2 = plt.subplot(gs[0, 2:4]) # Spans two columns
# Second and third rows: 4x2 grid (8 subplots)
ml_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2)]
gnn_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2, 4)]
ML_result = benchmark_results[benchmark_results["model_type"] == "Classical ML"]
GNN_result = benchmark_results[benchmark_results["model_type"] == "GNN"]
ML_result = ML_result[ML_result["split"] == split]
GNN_result = GNN_result[GNN_result["split"] == split]
# Plot main comparisons
make_regplot(ML_result[f"ID_test_{metric}"], ML_result[f"OOD_test_{metric}"], ax1)
make_regplot(GNN_result[f"ID_test_{metric}"], GNN_result[f"OOD_test_{metric}"], ax2, color="#8da0cb")
# Format main plots
for ax, title in [(ax1, f"Classical ML Models ({split})"), (ax2, f"GNN Models ({split})")]:
ax.set_title(title, fontsize=20, pad=10)
ax.set_xlabel(f"ID {metric_mapping[metric]}", fontsize=16)
ax.set_ylabel(f"OOD {metric_mapping[metric]}", fontsize=16)
if ax == ax2:
ax.set_ylabel("")
ax.tick_params(axis="both", labelsize=16)
format_axis(ax)
# For other axis, plot the same for ML and GNN for each datasets
# Plot and format datasets comparisons
for i, dataset in enumerate(DATASET_NAMES):
ML_dataset = ML_result[ML_result["dataset"] == dataset]
GNN_dataset = GNN_result[GNN_result["dataset"] == dataset]
# plot for the whole dataset in the background
sns.scatterplot(
x=f"ID_test_{metric}", y=f"OOD_test_{metric}", data=ML_result, ax=ml_axes[i], alpha=0.15, s=40, color="gray"
)
make_regplot(ML_dataset[f"ID_test_{metric}"], ML_dataset[f"OOD_test_{metric}"], ml_axes[i])
sns.scatterplot(
x=f"ID_test_{metric}", y=f"OOD_test_{metric}", data=GNN_result, ax=gnn_axes[i], alpha=0.15, s=40, color="gray"
)
make_regplot(GNN_dataset[f"ID_test_{metric}"], GNN_dataset[f"OOD_test_{metric}"], gnn_axes[i], color="#8da0cb")
# Format dataset plots
for ax in [ml_axes[i], gnn_axes[i]]:
ax.set_title(f"{dataset}", fontsize=18)
ax.tick_params(axis="both", labelsize=14)
format_axis(ax)
# Only show y-label for leftmost plots
if ax == ml_axes[i] and i % 2 == 0:
ax.set_ylabel(f"OOD {metric_mapping[metric]}", fontsize=12)
else:
ax.set_ylabel("")
# Only show x-label for bottom plots
if i >= len(SPLIT_TYPES) - 2:
ax.set_xlabel(f"ID {metric_mapping[metric]}", fontsize=12)
else:
ax.set_xlabel("")
# For the wole plots in row 1, add panel a, For the whole plots on next rows, add panel B
ax1.text(-0.15, 1.05, "(a)", transform=ax1.transAxes, fontsize=20, fontweight="bold")
ml_axes[0].text(-0.35, 1.15, "(b)", transform=ml_axes[0].transAxes, fontsize=20, fontweight="bold")
# Adjust layout to prevent label cutoff
# plt.tight_layout()
if save:
plt.savefig(f"assets/regplot_{split}_specific_with_categories.pdf", bbox_inches="tight")
plt.savefig(f"assets/regplot_{split}_with_categories.png", bbox_inches="tight", dpi=300)
plt.show()
# Choose oine specific splitting strategies. Then plot for all the datasets separately the relationship between ID and OOD test roc_auc
split = "max_dissimilarity"
save = False
metric = "roc_auc"
def make_regplot(x, y, ax, color="#66c2a5"):
sns.regplot(
x=x,
y=y,
ax=ax,
scatter_kws={"alpha": 0.5, "s": 40, "color": color},
line_kws={"color": "red", "linewidth": 2},
ci=95,
)
corr = pearsonr(x, y)[0]
props = dict(boxstyle="round", facecolor="white", alpha=0.8, edgecolor="gray")
ax.text(0.05, 0.95, f"r = {corr:.2f}", transform=ax.transAxes, fontsize=14, verticalalignment="top", bbox=props)
def format_axis(ax):
# lims = [
# np.min([ax.get_xlim(), ax.get_ylim()]),
# np.max([ax.get_xlim(), ax.get_ylim()])
# ]
lims = [0.5, 0.95]
ax.set_xlim(lims)
ax.set_ylim(lims)
ax.plot(lims, lims, "b--", alpha=0.5, zorder=0)
ax.set_xticks(np.arange(0.5, 1.0, 0.1))
ax.set_yticks(np.arange(0.5, 1.0, 0.1))
fig = plt.figure(figsize=(16, 12))
gs = gridspec.GridSpec(5, 4, height_ratios=[2, 1, 1, 1, 1], hspace=1.0, wspace=0.3) # 5 rows, 4 columns
# Add super title
fig.suptitle(f"ID vs OOD Performance Comparison ({metric_mapping[metric]})", fontsize=24, y=0.95)
# First row: One wide subplot
ax1 = plt.subplot(gs[0, 0:2]) # Spans two columns
ax2 = plt.subplot(gs[0, 2:4]) # Spans two columns
# Second and third rows: 4x2 grid (8 subplots)
ml_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2)]
gnn_axes = [plt.subplot(gs[i, j]) for i in range(1, 5) for j in range(2, 4)]
ML_result = benchmark_results[benchmark_results["model_type"] == "Classical ML"]
GNN_result = benchmark_results[benchmark_results["model_type"] == "GNN"]
ML_result = ML_result[ML_result["split"] == split]
GNN_result = GNN_result[GNN_result["split"] == split]
# Plot main comparisons
make_regplot(ML_result[f"ID_test_{metric}"], ML_result[f"OOD_test_{metric}"], ax1)
make_regplot(GNN_result[f"ID_test_{metric}"], GNN_result[f"OOD_test_{metric}"], ax2, color="#8da0cb")
# Format main plots
for ax, title in [(ax1, f"Classical ML Models ({split})"), (ax2, f"GNN Models ({split})")]:
ax.set_title(title, fontsize=20, pad=10)
ax.set_xlabel(f"ID {metric_mapping[metric]}", fontsize=16)
ax.set_ylabel(f"OOD {metric_mapping[metric]}", fontsize=16)
if ax == ax2:
ax.set_ylabel("")
ax.tick_params(axis="both", labelsize=16)
format_axis(ax)
# For other axis, plot the same for ML and GNN for each datasets
# Plot and format datasets comparisons
for i, dataset in enumerate(DATASET_NAMES):
ML_dataset = ML_result[ML_result["dataset"] == dataset]
GNN_dataset = GNN_result[GNN_result["dataset"] == dataset]
# plot for the whole dataset in the background
sns.scatterplot(
x=f"ID_test_{metric}", y=f"OOD_test_{metric}", data=ML_result, ax=ml_axes[i], alpha=0.15, s=40, color="gray"
)
make_regplot(ML_dataset[f"ID_test_{metric}"], ML_dataset[f"OOD_test_{metric}"], ml_axes[i])
sns.scatterplot(
x=f"ID_test_{metric}", y=f"OOD_test_{metric}", data=GNN_result, ax=gnn_axes[i], alpha=0.15, s=40, color="gray"
)
make_regplot(GNN_dataset[f"ID_test_{metric}"], GNN_dataset[f"OOD_test_{metric}"], gnn_axes[i], color="#8da0cb")
# Format dataset plots
for ax in [ml_axes[i], gnn_axes[i]]:
ax.set_title(f"{dataset}", fontsize=18)
ax.tick_params(axis="both", labelsize=14)
format_axis(ax)
# Only show y-label for leftmost plots
if ax == ml_axes[i] and i % 2 == 0:
ax.set_ylabel(f"OOD {metric_mapping[metric]}", fontsize=12)
else:
ax.set_ylabel("")
# Only show x-label for bottom plots
if i >= len(SPLIT_TYPES) - 2:
ax.set_xlabel(f"ID {metric_mapping[metric]}", fontsize=12)
else:
ax.set_xlabel("")
# For the wole plots in row 1, add panel a, For the whole plots on next rows, add panel B
ax1.text(-0.15, 1.05, "(a)", transform=ax1.transAxes, fontsize=20, fontweight="bold")
ml_axes[0].text(-0.35, 1.15, "(b)", transform=ml_axes[0].transAxes, fontsize=20, fontweight="bold")
# Adjust layout to prevent label cutoff
# plt.tight_layout()
if save:
plt.savefig(f"assets/regplot_{split}_specific_with_categories.pdf", bbox_inches="tight")
plt.savefig(f"assets/regplot_{split}_with_categories.png", bbox_inches="tight", dpi=300)
plt.show()
Datasets PhysicoChemical Properties (Figure S3)¶
In [ ]:
Copied!
# Function to compute physicochemical properties
def compute_properties(smiles):
mol = Chem.MolFromSmiles(smiles)
if mol:
return {
"Molecular_Weight": Descriptors.MolWt(mol),
"LogP": Descriptors.MolLogP(mol),
"TPSA": Descriptors.TPSA(mol),
"HBD": Descriptors.NumHDonors(mol),
"HBA": Descriptors.NumHAcceptors(mol),
"Rotatable_Bonds": Descriptors.NumRotatableBonds(mol),
}
else:
return {
"Molecular_Weight": None,
"LogP": None,
"TPSA": None,
"HBD": None,
"HBA": None,
"Rotatable_Bonds": None,
}
# Function to process a dataset
def process_dataset(df):
properties_df = df["smiles"].apply(compute_properties).apply(pd.Series)
return pd.concat([df, properties_df], axis=1)
# Function to compute physicochemical properties
def compute_properties(smiles):
mol = Chem.MolFromSmiles(smiles)
if mol:
return {
"Molecular_Weight": Descriptors.MolWt(mol),
"LogP": Descriptors.MolLogP(mol),
"TPSA": Descriptors.TPSA(mol),
"HBD": Descriptors.NumHDonors(mol),
"HBA": Descriptors.NumHAcceptors(mol),
"Rotatable_Bonds": Descriptors.NumRotatableBonds(mol),
}
else:
return {
"Molecular_Weight": None,
"LogP": None,
"TPSA": None,
"HBD": None,
"HBA": None,
"Rotatable_Bonds": None,
}
# Function to process a dataset
def process_dataset(df):
properties_df = df["smiles"].apply(compute_properties).apply(pd.Series)
return pd.concat([df, properties_df], axis=1)
In [ ]:
Copied!
props_df = []
for datasets in DATASET_NAMES:
df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", datasets, f"{datasets}_standardize.csv"))
processed_df = process_dataset(df)
props_df.append(processed_df)
props_df = []
for datasets in DATASET_NAMES:
df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", datasets, f"{datasets}_standardize.csv"))
processed_df = process_dataset(df)
props_df.append(processed_df)
In [ ]:
Copied!
def normalize_property(values, property_name):
"""
Normalize properties based on their characteristics
"""
if property_name == "LogP":
# Center around 0 with typical drug-like range (-2 to 5)
return (values - (-2)) / (5 - (-2))
elif property_name == "Molecular_Weight":
# Normalize based on typical drug-like range (160-480)
return (values - 160) / (480 - 160)
elif property_name == "TPSA":
# Normalize based on typical range (0-140)
return values / 140
elif property_name == "HBD":
# Normalize based on Lipinski's rule (≤5)
return values / 5
elif property_name == "HBA":
# Normalize based on Lipinski's rule (≤10)
return values / 10
elif property_name == "Rotatable_Bonds":
# Normalize based on typical flexibility rule (≤10)
return values / 10
else:
# Default min-max normalization
return (values - values.min()) / (values.max() - values.min())
def plot_radar_subplots(datasets, save=False):
properties = ["Molecular_Weight", "LogP", "TPSA", "HBD", "HBA", "Rotatable_Bonds"]
dataset_names = list(datasets.keys())
_ = len(dataset_names)
# Create figure with better spacing
fig = plt.figure(figsize=(12, 18))
gs = fig.add_gridspec(3, 2, hspace=0.1, wspace=0.2)
# Use a more distinguishable color palette
colors = sns.color_palette("Set2")
for i, property_name in enumerate(properties):
ax = fig.add_subplot(gs[i // 2, i % 2], projection="polar")
# Collect and normalize median values
values = [datasets[name][property_name].median() for name in dataset_names] # Changed to median
values = normalize_property(np.array(values), property_name)
# Calculate angles
angles = np.linspace(0, 2 * pi, len(dataset_names), endpoint=False)
# Close the radar chart
values = np.concatenate((values, [values[0]]))
angles = np.concatenate((angles, [angles[0]]))
# Plot with enhanced styling
ax.plot(angles, values, "o-", color=colors[i], linewidth=2.5, markersize=8)
ax.fill(angles, values, color=colors[i], alpha=0.25)
# Enhance grid and labels
ax.grid(True, color="gray", alpha=0.3, linewidth=0.5)
ax.set_xticks(angles[:-1])
# Rotate labels for better readability
ax.set_xticklabels([r"\textbf{" + name + "}" for name in dataset_names], fontsize=16)
# Add value markers at specific intervals
ax.set_yticks([0.2, 0.4, 0.6, 0.8, 1.0])
ax.set_yticklabels(["0.2", "0.4", "0.6", "0.8", "1.0"], fontsize=16)
# Add property statistics (using median and IQR)
median_val = np.median(values[:-1])
q1 = np.percentile(values[:-1], 25)
q3 = np.percentile(values[:-1], 75)
iqr = q3 - q1
stats_text = f"Median: {median_val:.2f}\nIQR: {iqr:.2f}"
ax.text(
1.0,
1.1,
stats_text,
transform=ax.transAxes,
bbox=dict(facecolor="white", alpha=0.8, edgecolor="none"),
fontsize=16,
ha="right",
va="top",
)
# Enhance title
title_text = rf"\textbf{{{property_name}}}" + "\n" + r"\textbf{Distribution (Median)}"
ax.set_title(title_text, fontsize=20, pad=22)
# Add main title
plt.suptitle(
r"\textbf{Physicochemical Properties Distribution Across Data Sets (Median Values)}",
fontsize=24,
fontweight="bold",
y=1.05,
)
# Add legend with property ranges
# For the legend text
legend_text = {
r"\textbf{Molecular Weight}": r"\textbf{Range: 160-480 Da}",
r"\textbf{LogP}": r"\textbf{Range: -2 to 5}",
r"\textbf{TPSA}": r"\textbf{Range: 0-140 Ų}",
r"\textbf{HBD}": r"\textbf{Range: 0-5}",
r"\textbf{HBA}": r"\textbf{Range: 0-10}",
r"\textbf{Rotatable Bonds}": r"\textbf{Range: 0-10}",
}
fig.text(
1.02,
0.9,
"\n".join([f"{k}: {v}" for k, v in legend_text.items()]),
fontsize=18,
transform=fig.transFigure,
bbox=dict(facecolor="white", alpha=0.8, edgecolor="none", pad=12),
)
if save:
plt.savefig("assets/figures/radar_subplots.pdf", bbox_inches="tight")
plt.savefig("assets/figures/radar_subplots.png", bbox_inches="tight", dpi=600)
plt.show()
datasets = {}
for i, dataset in enumerate(DATASET_NAMES):
datasets[dataset] = props_df[i]
# Use the function
plot_radar_subplots(datasets, save=True)
def normalize_property(values, property_name):
"""
Normalize properties based on their characteristics
"""
if property_name == "LogP":
# Center around 0 with typical drug-like range (-2 to 5)
return (values - (-2)) / (5 - (-2))
elif property_name == "Molecular_Weight":
# Normalize based on typical drug-like range (160-480)
return (values - 160) / (480 - 160)
elif property_name == "TPSA":
# Normalize based on typical range (0-140)
return values / 140
elif property_name == "HBD":
# Normalize based on Lipinski's rule (≤5)
return values / 5
elif property_name == "HBA":
# Normalize based on Lipinski's rule (≤10)
return values / 10
elif property_name == "Rotatable_Bonds":
# Normalize based on typical flexibility rule (≤10)
return values / 10
else:
# Default min-max normalization
return (values - values.min()) / (values.max() - values.min())
def plot_radar_subplots(datasets, save=False):
properties = ["Molecular_Weight", "LogP", "TPSA", "HBD", "HBA", "Rotatable_Bonds"]
dataset_names = list(datasets.keys())
_ = len(dataset_names)
# Create figure with better spacing
fig = plt.figure(figsize=(12, 18))
gs = fig.add_gridspec(3, 2, hspace=0.1, wspace=0.2)
# Use a more distinguishable color palette
colors = sns.color_palette("Set2")
for i, property_name in enumerate(properties):
ax = fig.add_subplot(gs[i // 2, i % 2], projection="polar")
# Collect and normalize median values
values = [datasets[name][property_name].median() for name in dataset_names] # Changed to median
values = normalize_property(np.array(values), property_name)
# Calculate angles
angles = np.linspace(0, 2 * pi, len(dataset_names), endpoint=False)
# Close the radar chart
values = np.concatenate((values, [values[0]]))
angles = np.concatenate((angles, [angles[0]]))
# Plot with enhanced styling
ax.plot(angles, values, "o-", color=colors[i], linewidth=2.5, markersize=8)
ax.fill(angles, values, color=colors[i], alpha=0.25)
# Enhance grid and labels
ax.grid(True, color="gray", alpha=0.3, linewidth=0.5)
ax.set_xticks(angles[:-1])
# Rotate labels for better readability
ax.set_xticklabels([r"\textbf{" + name + "}" for name in dataset_names], fontsize=16)
# Add value markers at specific intervals
ax.set_yticks([0.2, 0.4, 0.6, 0.8, 1.0])
ax.set_yticklabels(["0.2", "0.4", "0.6", "0.8", "1.0"], fontsize=16)
# Add property statistics (using median and IQR)
median_val = np.median(values[:-1])
q1 = np.percentile(values[:-1], 25)
q3 = np.percentile(values[:-1], 75)
iqr = q3 - q1
stats_text = f"Median: {median_val:.2f}\nIQR: {iqr:.2f}"
ax.text(
1.0,
1.1,
stats_text,
transform=ax.transAxes,
bbox=dict(facecolor="white", alpha=0.8, edgecolor="none"),
fontsize=16,
ha="right",
va="top",
)
# Enhance title
title_text = rf"\textbf{{{property_name}}}" + "\n" + r"\textbf{Distribution (Median)}"
ax.set_title(title_text, fontsize=20, pad=22)
# Add main title
plt.suptitle(
r"\textbf{Physicochemical Properties Distribution Across Data Sets (Median Values)}",
fontsize=24,
fontweight="bold",
y=1.05,
)
# Add legend with property ranges
# For the legend text
legend_text = {
r"\textbf{Molecular Weight}": r"\textbf{Range: 160-480 Da}",
r"\textbf{LogP}": r"\textbf{Range: -2 to 5}",
r"\textbf{TPSA}": r"\textbf{Range: 0-140 Ų}",
r"\textbf{HBD}": r"\textbf{Range: 0-5}",
r"\textbf{HBA}": r"\textbf{Range: 0-10}",
r"\textbf{Rotatable Bonds}": r"\textbf{Range: 0-10}",
}
fig.text(
1.02,
0.9,
"\n".join([f"{k}: {v}" for k, v in legend_text.items()]),
fontsize=18,
transform=fig.transFigure,
bbox=dict(facecolor="white", alpha=0.8, edgecolor="none", pad=12),
)
if save:
plt.savefig("assets/figures/radar_subplots.pdf", bbox_inches="tight")
plt.savefig("assets/figures/radar_subplots.png", bbox_inches="tight", dpi=600)
plt.show()
datasets = {}
for i, dataset in enumerate(DATASET_NAMES):
datasets[dataset] = props_df[i]
# Use the function
plot_radar_subplots(datasets, save=True)
Datasets Size Comparisson For Different Splitters (ID vs OOD test size) (Figure S4)¶
In [ ]:
Copied!
# Fixes split type. Then, find out ratio of test set size to the size of the dataset
dataset_category = "TDC"
dataset_names = DATASET_NAMES
# dataset_names = "CYP1A2"
split_types = SPLIT_TYPES
dfs = []
for dataset_name in dataset_names:
for split_type in split_types:
dataset_folder = os.path.join(DATASET_PATH, dataset_category, dataset_name, "split", split_type)
id_size = []
id_test_size = []
ood_test_size = []
df = pd.DataFrame()
with open(os.path.join(dataset_folder, "config.json"), "r") as f:
data_config = json.load(f)
for i in range(10):
id_size.append(data_config[f"train_size_{i}"])
ood_test_size.append(data_config[f"test_size_{i}"])
id_test_size.append(data_config[f"train_size_{i}"] * 0.2)
id_frac = [id_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
ood_test_frac = [ood_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
id_test_frac = [id_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
df["ood_test_size"] = ood_test_frac
df["id_test_size"] = id_test_frac
df["split_type"] = [split_type] * 10
df["dataset"] = [dataset_name] * 10
dfs.append(df)
df = pd.concat(dfs)
# Fixes split type. Then, find out ratio of test set size to the size of the dataset
dataset_category = "TDC"
dataset_names = DATASET_NAMES
# dataset_names = "CYP1A2"
split_types = SPLIT_TYPES
dfs = []
for dataset_name in dataset_names:
for split_type in split_types:
dataset_folder = os.path.join(DATASET_PATH, dataset_category, dataset_name, "split", split_type)
id_size = []
id_test_size = []
ood_test_size = []
df = pd.DataFrame()
with open(os.path.join(dataset_folder, "config.json"), "r") as f:
data_config = json.load(f)
for i in range(10):
id_size.append(data_config[f"train_size_{i}"])
ood_test_size.append(data_config[f"test_size_{i}"])
id_test_size.append(data_config[f"train_size_{i}"] * 0.2)
id_frac = [id_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
ood_test_frac = [ood_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
id_test_frac = [id_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
df["ood_test_size"] = ood_test_frac
df["id_test_size"] = id_test_frac
df["split_type"] = [split_type] * 10
df["dataset"] = [dataset_name] * 10
dfs.append(df)
df = pd.concat(dfs)
In [ ]:
Copied!
# boxplot of test size ratio for split types (aggregate over datasets)
save = True
# Create figure with higher DPI and adjusted size for publication
fig, ax = plt.subplots(figsize=(10, 6), dpi=300)
# Melt the dataframe
df_melt = df.melt(id_vars=["split_type", "dataset"], var_name="test_type", value_name="size_ratio")
# Create boxplot with publication-ready styling
sns.boxplot(
x="split_type",
y="size_ratio",
hue="test_type",
data=df_melt,
ax=ax,
palette="Set2",
linewidth=1.5,
fliersize=2,
width=0.7,
capwidths=0.1,
)
# Customize axes
ax.set_ylabel(r"\textbf{Test Set Size (\%)}", fontsize=18, fontweight="bold")
ax.set_xlabel(r"\textbf{Split Type}", fontsize=18, fontweight="bold")
ax.set_title(r"\textbf{Distribution of Test Set Sizes Across Different Split Types}", fontsize=20, pad=20)
# x-ticks should be split type mapping
ax.set_xticks(range(len(SPLIT_TYPPE_MAPPING)))
ax.set_xticklabels(
[r"\textbf{" + splitter + "}" for splitter in SPLIT_TYPPE_MAPPING.values()], rotation=45, ha="right", fontsize=14
)
# Enhance grid
ax.grid(axis="y", linestyle="--", alpha=0.3, color="gray")
ax.set_axisbelow(True)
# Customize legend
handles, labels = ax.get_legend_handles_labels()
ax.legend(
handles=handles,
labels=["Out-of-Distribution", "In-Distribution"],
title="Test Set Type",
title_fontsize=12,
fontsize=11,
loc="upper right",
frameon=True,
framealpha=0.7,
edgecolor="gray",
facecolor="white",
)
# Adjust tick labels
plt.xticks(rotation=45, ha="right", fontsize=11)
plt.yticks(fontsize=11)
# Save with high quality
if save:
fig.savefig("assets/figures/test_size_ratio.pdf", bbox_inches="tight", dpi=600)
fig.savefig("assets/figures/test_size_ratio.png", bbox_inches="tight", dpi=600)
plt.show()
# boxplot of test size ratio for split types (aggregate over datasets)
save = True
# Create figure with higher DPI and adjusted size for publication
fig, ax = plt.subplots(figsize=(10, 6), dpi=300)
# Melt the dataframe
df_melt = df.melt(id_vars=["split_type", "dataset"], var_name="test_type", value_name="size_ratio")
# Create boxplot with publication-ready styling
sns.boxplot(
x="split_type",
y="size_ratio",
hue="test_type",
data=df_melt,
ax=ax,
palette="Set2",
linewidth=1.5,
fliersize=2,
width=0.7,
capwidths=0.1,
)
# Customize axes
ax.set_ylabel(r"\textbf{Test Set Size (\%)}", fontsize=18, fontweight="bold")
ax.set_xlabel(r"\textbf{Split Type}", fontsize=18, fontweight="bold")
ax.set_title(r"\textbf{Distribution of Test Set Sizes Across Different Split Types}", fontsize=20, pad=20)
# x-ticks should be split type mapping
ax.set_xticks(range(len(SPLIT_TYPPE_MAPPING)))
ax.set_xticklabels(
[r"\textbf{" + splitter + "}" for splitter in SPLIT_TYPPE_MAPPING.values()], rotation=45, ha="right", fontsize=14
)
# Enhance grid
ax.grid(axis="y", linestyle="--", alpha=0.3, color="gray")
ax.set_axisbelow(True)
# Customize legend
handles, labels = ax.get_legend_handles_labels()
ax.legend(
handles=handles,
labels=["Out-of-Distribution", "In-Distribution"],
title="Test Set Type",
title_fontsize=12,
fontsize=11,
loc="upper right",
frameon=True,
framealpha=0.7,
edgecolor="gray",
facecolor="white",
)
# Adjust tick labels
plt.xticks(rotation=45, ha="right", fontsize=11)
plt.yticks(fontsize=11)
# Save with high quality
if save:
fig.savefig("assets/figures/test_size_ratio.pdf", bbox_inches="tight", dpi=600)
fig.savefig("assets/figures/test_size_ratio.png", bbox_inches="tight", dpi=600)
plt.show()
HIT Rate Plots¶
Hit Rate vs ROC-AUC¶
In [ ]:
Copied!
# Load hit rate data and plot scatter plots for each splitter
save = True
# Load hit rate data
hit_rate_df = pd.read_csv("classification_results/TDC/hit_rate.csv")
# Get unique splitters
splitters = hit_rate_df["splitter"].unique()
# Create subplot for each splitter
fig, axes = plt.subplots(3, 4, figsize=(20, 10), gridspec_kw={"wspace": 0.4, "hspace": 0.7})
axes = axes.flatten()
for i, splitter in enumerate(splitters):
splitter_data = hit_rate_df[hit_rate_df["splitter"] == splitter]
# Create scatter plot
sns.scatterplot(data=splitter_data, x="roc_auc", y="hit_rate", ax=axes[i], alpha=0.6, s=40)
# Add y=x line
axes[i].plot([0.4, 1.0], [0.4, 1.0], "--", color="gray", alpha=1)
# Calculate correlations
pearson_r = splitter_data["roc_auc"].corr(splitter_data["hit_rate"], method="pearson")
spearman_r = splitter_data["roc_auc"].corr(splitter_data["hit_rate"], method="spearman")
# Add correlation text
axes[i].text(
0.72,
0.1,
f"Pearson's r: {pearson_r:.2f}\nSpearman's r: {spearman_r:.2f}",
fontsize=12,
bbox=dict(facecolor="white", alpha=0.5, edgecolor="gray", pad=3),
)
# Format axis labels and title using mapping
axes[i].set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=16, pad=10)
axes[i].set_xlabel(r"$\textbf{ROC-AUC}$", fontsize=12, labelpad=10)
axes[i].set_ylabel(r"$\textbf{Hit Rate}$", fontsize=12, labelpad=10)
# Set axis limits consistently
axes[i].set_xlim(0.4, 1.0)
axes[i].set_ylim(0, 1.15)
# Add grid
axes[i].grid(True, linestyle="-", alpha=0.7)
# Remove top and right spines
axes[i].spines["top"].set_visible(False)
axes[i].spines["right"].set_visible(False)
# Remove the empty 12th subplot
axes[-1].remove()
# Adjust layout
plt.tight_layout()
# Add a main title
fig.suptitle(r"$\textbf{Hit Rate vs ROC-AUC for Different Splitting Methods}$", fontsize=20, y=0.97)
if save:
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "hit_rate_vs_roc_auc_TDC.pdf"), bbox_inches="tight", dpi=600
)
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "hit_rate_vs_roc_auc_TDC.png"), bbox_inches="tight", dpi=600
)
plt.show()
# Load hit rate data and plot scatter plots for each splitter
save = True
# Load hit rate data
hit_rate_df = pd.read_csv("classification_results/TDC/hit_rate.csv")
# Get unique splitters
splitters = hit_rate_df["splitter"].unique()
# Create subplot for each splitter
fig, axes = plt.subplots(3, 4, figsize=(20, 10), gridspec_kw={"wspace": 0.4, "hspace": 0.7})
axes = axes.flatten()
for i, splitter in enumerate(splitters):
splitter_data = hit_rate_df[hit_rate_df["splitter"] == splitter]
# Create scatter plot
sns.scatterplot(data=splitter_data, x="roc_auc", y="hit_rate", ax=axes[i], alpha=0.6, s=40)
# Add y=x line
axes[i].plot([0.4, 1.0], [0.4, 1.0], "--", color="gray", alpha=1)
# Calculate correlations
pearson_r = splitter_data["roc_auc"].corr(splitter_data["hit_rate"], method="pearson")
spearman_r = splitter_data["roc_auc"].corr(splitter_data["hit_rate"], method="spearman")
# Add correlation text
axes[i].text(
0.72,
0.1,
f"Pearson's r: {pearson_r:.2f}\nSpearman's r: {spearman_r:.2f}",
fontsize=12,
bbox=dict(facecolor="white", alpha=0.5, edgecolor="gray", pad=3),
)
# Format axis labels and title using mapping
axes[i].set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=16, pad=10)
axes[i].set_xlabel(r"$\textbf{ROC-AUC}$", fontsize=12, labelpad=10)
axes[i].set_ylabel(r"$\textbf{Hit Rate}$", fontsize=12, labelpad=10)
# Set axis limits consistently
axes[i].set_xlim(0.4, 1.0)
axes[i].set_ylim(0, 1.15)
# Add grid
axes[i].grid(True, linestyle="-", alpha=0.7)
# Remove top and right spines
axes[i].spines["top"].set_visible(False)
axes[i].spines["right"].set_visible(False)
# Remove the empty 12th subplot
axes[-1].remove()
# Adjust layout
plt.tight_layout()
# Add a main title
fig.suptitle(r"$\textbf{Hit Rate vs ROC-AUC for Different Splitting Methods}$", fontsize=20, y=0.97)
if save:
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "hit_rate_vs_roc_auc_TDC.pdf"), bbox_inches="tight", dpi=600
)
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "hit_rate_vs_roc_auc_TDC.png"), bbox_inches="tight", dpi=600
)
plt.show()
Hit Rate Report For All Splitters¶
In [ ]:
Copied!
# Create grouped bar plot for ML vs GNN model differences
save = True
fig = plt.figure(figsize=(16, 12))
# gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, wspace=0.5, figure=fig) # Increased hspace for more gap
gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, figure=fig) # Increased hspace for more gap
# First row: one wide subplot
gs1 = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=gs[0])
ax1 = plt.subplot(gs1[0, 0])
# Second and third rows: 4x2 grid (8 subplots)
gs2 = gridspec.GridSpecFromSubplotSpec(2, 4, subplot_spec=gs[1])
axes = [plt.subplot(gs2[i, j]) for i in range(2) for j in range(4)]
df = hit_rate.groupby(["splitter", "model_type", "dataset"])["hit_rate"].mean().reset_index()
df["splitter"] = df["splitter"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["splitter"] = pd.Categorical(df["splitter"], split_order)
# Main plot
g = sns.barplot(x="splitter", y="hit_rate", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_xlabel("") # No x-label on main plot
ax1.set_ylabel(r"\textbf{Hit Rate}", fontsize=20)
ax1.set_title(r"\textbf{Hit Rate Distribution Across Splitting Methods}", fontsize=24, pad=20)
# ticks formatting
format_xticklabels(ax1, rotation=30, ha="right", fontsize=16)
plt.setp(ax1.get_yticklabels(), fontsize=16)
# Grid and legend
ax1.grid(axis="y", linestyle="--", alpha=0.6)
legend = ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=14, title=r"\textbf{Model Type}")
for text in legend.get_texts():
text.set_fontsize(14)
text.set_text(r"\textbf{" + text.get_text() + "}")
# Subplots
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="splitter", y="hit_rate", hue="model_type", data=df_dataset, ax=axes[i])
axes[i].set_title(r"\textbf{" + dataset + "}", fontsize=18, pad=10)
axes[i].set_xlabel("")
if i % 4 == 0:
axes[i].set_ylabel(r"\textbf{Hit Rate}", fontsize=14)
else:
axes[i].set_ylabel("")
# x-tick formatting
if i >= 4:
format_xticklabels(axes[i])
else:
axes[i].set_xticklabels([])
# y-tick formatting
# yticklabels = [label.get_text() for label in axes[i].get_yticklabels()]
# axes[i].set_yticklabels([r'\textbf{' + label + '}' for label in yticklabels], fontsize=10)
axes[i].tick_params(axis="both", labelsize=14)
axes[i].grid(axis="y", linestyle="--", alpha=0.5)
g.legend().remove()
axes[i].set_ylim(0.3, 1.0)
axes[i].set_yticks(np.arange(0.3, 1.0, 0.1))
# Subplot labels
ax1.text(-0.08, 1.15, r"\textbf{(a)}", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.27, 1.35, r"\textbf{(b)}", transform=axes[0].transAxes, fontsize=20)
# plt.subplots_adjust(left=0.07, right=0.95, top=0.93, bottom=0.08, hspace=1, wspace=0.45) # Increased hspace
if save:
plt.savefig("assets/figures/grouped_barplot_hit_rate.pdf", bbox_inches="tight")
plt.savefig("assets/figures/grouped_barplot_hit_rate.png", bbox_inches="tight", dpi=600) # High-res
plt.show()
# Create grouped bar plot for ML vs GNN model differences
save = True
fig = plt.figure(figsize=(16, 12))
# gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, wspace=0.5, figure=fig) # Increased hspace for more gap
gs = gridspec.GridSpec(2, 1, height_ratios=[1, 1], hspace=0.8, figure=fig) # Increased hspace for more gap
# First row: one wide subplot
gs1 = gridspec.GridSpecFromSubplotSpec(1, 1, subplot_spec=gs[0])
ax1 = plt.subplot(gs1[0, 0])
# Second and third rows: 4x2 grid (8 subplots)
gs2 = gridspec.GridSpecFromSubplotSpec(2, 4, subplot_spec=gs[1])
axes = [plt.subplot(gs2[i, j]) for i in range(2) for j in range(4)]
df = hit_rate.groupby(["splitter", "model_type", "dataset"])["hit_rate"].mean().reset_index()
df["splitter"] = df["splitter"].map(SPLIT_TYPPE_MAPPING)
# rearrange the order of the split
split_order = [SPLIT_TYPPE_MAPPING[split] for split in SPLIT_TYPES]
df["splitter"] = pd.Categorical(df["splitter"], split_order)
# Main plot
g = sns.barplot(x="splitter", y="hit_rate", hue="model_type", data=df, ax=ax1, legend=True)
ax1.set_xlabel("") # No x-label on main plot
ax1.set_ylabel(r"\textbf{Hit Rate}", fontsize=20)
ax1.set_title(r"\textbf{Hit Rate Distribution Across Splitting Methods}", fontsize=24, pad=20)
# ticks formatting
format_xticklabels(ax1, rotation=30, ha="right", fontsize=16)
plt.setp(ax1.get_yticklabels(), fontsize=16)
# Grid and legend
ax1.grid(axis="y", linestyle="--", alpha=0.6)
legend = ax1.legend(loc="upper left", bbox_to_anchor=(1, 1), fontsize=14, title=r"\textbf{Model Type}")
for text in legend.get_texts():
text.set_fontsize(14)
text.set_text(r"\textbf{" + text.get_text() + "}")
# Subplots
for i, dataset in enumerate(dataset_names):
df_dataset = df[df["dataset"] == dataset]
g = sns.barplot(x="splitter", y="hit_rate", hue="model_type", data=df_dataset, ax=axes[i])
axes[i].set_title(r"\textbf{" + dataset + "}", fontsize=18, pad=10)
axes[i].set_xlabel("")
if i % 4 == 0:
axes[i].set_ylabel(r"\textbf{Hit Rate}", fontsize=14)
else:
axes[i].set_ylabel("")
# x-tick formatting
if i >= 4:
format_xticklabels(axes[i])
else:
axes[i].set_xticklabels([])
# y-tick formatting
# yticklabels = [label.get_text() for label in axes[i].get_yticklabels()]
# axes[i].set_yticklabels([r'\textbf{' + label + '}' for label in yticklabels], fontsize=10)
axes[i].tick_params(axis="both", labelsize=14)
axes[i].grid(axis="y", linestyle="--", alpha=0.5)
g.legend().remove()
axes[i].set_ylim(0.3, 1.0)
axes[i].set_yticks(np.arange(0.3, 1.0, 0.1))
# Subplot labels
ax1.text(-0.08, 1.15, r"\textbf{(a)}", transform=ax1.transAxes, fontsize=20)
axes[0].text(-0.27, 1.35, r"\textbf{(b)}", transform=axes[0].transAxes, fontsize=20)
# plt.subplots_adjust(left=0.07, right=0.95, top=0.93, bottom=0.08, hspace=1, wspace=0.45) # Increased hspace
if save:
plt.savefig("assets/figures/grouped_barplot_hit_rate.pdf", bbox_inches="tight")
plt.savefig("assets/figures/grouped_barplot_hit_rate.png", bbox_inches="tight", dpi=600) # High-res
plt.show()
TSNE Plot for one Dataset and all Splitters¶
In [ ]:
Copied!
# Dataset path
dataset = "CYP2C19"
save = True
# load it back
data = np.load(f"datasets/TDC/{dataset}/features_2d.npz")
features_2d = data["features_2d"]
# Load the original dataset features and calculate Morgan fingerprints
data_df = pd.read_csv(f"datasets/TDC/{dataset}/{dataset}_standardize.csv")
# Create figure with 11 subplots (one for each splitter)
fig, axes = plt.subplots(3, 4, figsize=(16, 12), gridspec_kw={"wspace": 0.2, "hspace": 0.4})
axes = axes.ravel()
splits = SPLIT_TYPES
# splits =["random"]
# Iterate through each splitter
for i, splitter in enumerate(splits):
# Load train and test sets for fold 5
train = pd.read_csv(f"datasets/TDC/{dataset}/split/{splitter}/train_2.csv")
test = pd.read_csv(f"datasets/TDC/{dataset}/split/{splitter}/test_2.csv")
# find out exactly the index that smiles of train and test csv overlap with standadize data
train_index = [data_df.index[data_df["smiles"] == smiles].tolist() for smiles in train["smiles"]]
train_index = [idx for sublist in train_index for idx in sublist] # Flatten the list
test_index = [data_df.index[data_df["smiles"] == smiles].tolist() for smiles in test["smiles"]]
test_index = [idx for sublist in test_index for idx in sublist] # Flatten the list
# Create boolean masks for train and test
train_mask = np.zeros(len(features_2d), dtype=bool)
test_mask = np.zeros(len(features_2d), dtype=bool)
train_mask[train_index] = True
test_mask[test_index] = True
# Plot training set first (background)
scatter_train = axes[i].scatter(
features_2d[train_mask, 0], features_2d[train_mask, 1], c="lightblue", alpha=0.8, s=20, label="ID Data Set"
)
# Plot test set second (foreground, on top)
scatter_test = axes[i].scatter(
features_2d[test_mask, 0], features_2d[test_mask, 1], c="orange", alpha=0.6, s=5, label="OOD Test Set"
)
# Add title and format plot
axes[i].set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=16, pad=10)
axes[i].set_xticks([])
axes[i].set_yticks([])
axes[i].spines["top"].set_visible(False)
axes[i].spines["right"].set_visible(False)
axes[i].spines["left"].set_visible(False)
axes[i].spines["bottom"].set_visible(False)
# Remove the empty 12th subplot
axes[-1].remove()
# Add colorbar
# cbar = plt.colorbar(scatter, ax=axes, orientation='vertical', pad=0.02)
# cbar.set_ticks([0, 1])
# cbar.set_ticklabels(['Test Set', 'Train Set'])
# get legend, put it to the bottom right of the figure
handles, labels = axes[0].get_legend_handles_labels()
legend = fig.legend(
handles=handles, labels=labels, loc="lower right", bbox_to_anchor=(0.85, 0.2), title=r"\textbf{Label}", fontsize=16
)
# Adjust layout
plt.tight_layout()
# Add main title
plt.suptitle(
rf"$\textbf{{t-SNE Visualization of Different Splitting Methods}}$ ({dataset}, Fold 5)", fontsize=20, y=0.95
)
if save:
plt.savefig(os.path.join(CHECKOUT_PATH, "assets", "figures", "tSNE_visualization_TDC.pdf"), bbox_inches="tight")
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "tSNE_visualization_TDC.png"), bbox_inches="tight", dpi=600
)
plt.show()
# Dataset path
dataset = "CYP2C19"
save = True
# load it back
data = np.load(f"datasets/TDC/{dataset}/features_2d.npz")
features_2d = data["features_2d"]
# Load the original dataset features and calculate Morgan fingerprints
data_df = pd.read_csv(f"datasets/TDC/{dataset}/{dataset}_standardize.csv")
# Create figure with 11 subplots (one for each splitter)
fig, axes = plt.subplots(3, 4, figsize=(16, 12), gridspec_kw={"wspace": 0.2, "hspace": 0.4})
axes = axes.ravel()
splits = SPLIT_TYPES
# splits =["random"]
# Iterate through each splitter
for i, splitter in enumerate(splits):
# Load train and test sets for fold 5
train = pd.read_csv(f"datasets/TDC/{dataset}/split/{splitter}/train_2.csv")
test = pd.read_csv(f"datasets/TDC/{dataset}/split/{splitter}/test_2.csv")
# find out exactly the index that smiles of train and test csv overlap with standadize data
train_index = [data_df.index[data_df["smiles"] == smiles].tolist() for smiles in train["smiles"]]
train_index = [idx for sublist in train_index for idx in sublist] # Flatten the list
test_index = [data_df.index[data_df["smiles"] == smiles].tolist() for smiles in test["smiles"]]
test_index = [idx for sublist in test_index for idx in sublist] # Flatten the list
# Create boolean masks for train and test
train_mask = np.zeros(len(features_2d), dtype=bool)
test_mask = np.zeros(len(features_2d), dtype=bool)
train_mask[train_index] = True
test_mask[test_index] = True
# Plot training set first (background)
scatter_train = axes[i].scatter(
features_2d[train_mask, 0], features_2d[train_mask, 1], c="lightblue", alpha=0.8, s=20, label="ID Data Set"
)
# Plot test set second (foreground, on top)
scatter_test = axes[i].scatter(
features_2d[test_mask, 0], features_2d[test_mask, 1], c="orange", alpha=0.6, s=5, label="OOD Test Set"
)
# Add title and format plot
axes[i].set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=16, pad=10)
axes[i].set_xticks([])
axes[i].set_yticks([])
axes[i].spines["top"].set_visible(False)
axes[i].spines["right"].set_visible(False)
axes[i].spines["left"].set_visible(False)
axes[i].spines["bottom"].set_visible(False)
# Remove the empty 12th subplot
axes[-1].remove()
# Add colorbar
# cbar = plt.colorbar(scatter, ax=axes, orientation='vertical', pad=0.02)
# cbar.set_ticks([0, 1])
# cbar.set_ticklabels(['Test Set', 'Train Set'])
# get legend, put it to the bottom right of the figure
handles, labels = axes[0].get_legend_handles_labels()
legend = fig.legend(
handles=handles, labels=labels, loc="lower right", bbox_to_anchor=(0.85, 0.2), title=r"\textbf{Label}", fontsize=16
)
# Adjust layout
plt.tight_layout()
# Add main title
plt.suptitle(
rf"$\textbf{{t-SNE Visualization of Different Splitting Methods}}$ ({dataset}, Fold 5)", fontsize=20, y=0.95
)
if save:
plt.savefig(os.path.join(CHECKOUT_PATH, "assets", "figures", "tSNE_visualization_TDC.pdf"), bbox_inches="tight")
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "tSNE_visualization_TDC.png"), bbox_inches="tight", dpi=600
)
plt.show()
ID-OOD From Model Perspective¶
In [ ]:
Copied!
axes[i].get_xticklabels()
axes[i].get_xticklabels()
In [ ]:
Copied!
# Create plot with subplots for each splitter
save = True
sns.set_palette("Set2")
agg_func = "mean"
# Function to sort df_means based on MODEL_MAPPING order
def sort_df_by_model_mapping(df_means, model_mapping):
"""
Sort dataframe by the order of models in MODEL_MAPPING dictionary
Args:
df_means: DataFrame with 'model' column
model_mapping: Dictionary with model order
Returns:
Sorted DataFrame
"""
# Create a mapping of model names to their order in MODEL_MAPPING
model_order = {model: idx for idx, model in enumerate(model_mapping.keys())}
# Add a temporary column for sorting
df_means["_sort_order"] = df_means["model"].map(model_order)
# Sort by the order and drop the temporary column
df_sorted = df_means.sort_values("_sort_order").drop("_sort_order", axis=1).reset_index(drop=True)
return df_sorted
# Set up the figure and grid
fig = plt.figure(figsize=(22, 18))
gs = gridspec.GridSpec(4, 3, figure=fig, hspace=0.15, wspace=0.15)
axes = []
for i in range(11):
row = i // 3
col = i % 3
axes.append(plt.subplot(gs[row, col]))
# Get unique splitters
splitters = benchmark_results["split"].unique()
# For each splitter, create a subplot
for i, splitter in enumerate(splitters):
# Filter data for this splitter
df_split = benchmark_results[benchmark_results["split"] == splitter]
# Calculate mean values for each model
df_stats = df_split.groupby("model").agg({"ID_test_roc_auc": agg_func, "OOD_test_roc_auc": agg_func}).reset_index()
# Flatten column names
df_stats.columns = ["model", f"ID_{agg_func}", f"OOD_{agg_func}"]
# Calculate gap between ID and OOD
df_stats["gap"] = df_stats[f"ID_{agg_func}"] - df_stats[f"OOD_{agg_func}"]
# Sort based on MODEL_MAPPING order
df_stats = sort_df_by_model_mapping(df_stats, MODEL_MAPPING)
# Set up bar positions
x = np.arange(len(df_stats))
width = 0.35
# Create bars
axes[i].bar(x - width / 2, df_stats[f"ID_{agg_func}"], width, label="ID Test")
axes[i].bar(x + width / 2, df_stats[f"OOD_{agg_func}"], width, label="OOD Test")
# Add gap values as text above bars
for j, gap in enumerate(df_stats["gap"]):
# Position text above the higher of the two bars
y_pos = max(df_stats[f"ID_{agg_func}"].iloc[j], df_stats[f"OOD_{agg_func}"].iloc[j])
gap_text = f"{gap:.3f}"
axes[i].text(x[j], y_pos + 0.01, gap_text, ha="center", va="bottom", fontsize=10)
# Customize plot
axes[i].set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=22, pad=10)
axes[i].set_xticks(x)
axes[i].set_xticklabels(
[r"\textbf{" + MODEL_MAPPING[model] + "}" for model in df_stats["model"]], rotation=45, ha="right", fontsize=14
)
axes[i].set_ylabel(r"\textbf{ROC AUC}", fontsize=16, labelpad=10)
axes[i].grid(True, axis="y", linestyle="--", alpha=0.7) # Only show y-axis grid
axes[i].grid(False, axis="x") # Explicitly turn off x-axis grid
axes[i].set_ylim(0.5, 0.9) # Adjust as needed
# get legend, put it to the bottom right of the figure
handles, labels = axes[0].get_legend_handles_labels()
legend = fig.legend(
handles=handles,
labels=labels,
loc="lower right",
bbox_to_anchor=(0.8, 0.13),
title=r"\textbf{Test set type}",
fontsize=20,
title_fontsize=22,
)
# Remove any unused subplots
if len(axes) > len(splitters):
for j in range(len(splitters), len(axes)):
axes[j].remove()
plt.suptitle(r"$\textbf{Model Performance Across Different Splitting Methods}$", fontsize=26, y=1.05)
if save:
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.pdf"), bbox_inches="tight"
)
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.png"),
bbox_inches="tight",
dpi=600,
)
plt.show()
# Create plot with subplots for each splitter
save = True
sns.set_palette("Set2")
agg_func = "mean"
# Function to sort df_means based on MODEL_MAPPING order
def sort_df_by_model_mapping(df_means, model_mapping):
"""
Sort dataframe by the order of models in MODEL_MAPPING dictionary
Args:
df_means: DataFrame with 'model' column
model_mapping: Dictionary with model order
Returns:
Sorted DataFrame
"""
# Create a mapping of model names to their order in MODEL_MAPPING
model_order = {model: idx for idx, model in enumerate(model_mapping.keys())}
# Add a temporary column for sorting
df_means["_sort_order"] = df_means["model"].map(model_order)
# Sort by the order and drop the temporary column
df_sorted = df_means.sort_values("_sort_order").drop("_sort_order", axis=1).reset_index(drop=True)
return df_sorted
# Set up the figure and grid
fig = plt.figure(figsize=(22, 18))
gs = gridspec.GridSpec(4, 3, figure=fig, hspace=0.15, wspace=0.15)
axes = []
for i in range(11):
row = i // 3
col = i % 3
axes.append(plt.subplot(gs[row, col]))
# Get unique splitters
splitters = benchmark_results["split"].unique()
# For each splitter, create a subplot
for i, splitter in enumerate(splitters):
# Filter data for this splitter
df_split = benchmark_results[benchmark_results["split"] == splitter]
# Calculate mean values for each model
df_stats = df_split.groupby("model").agg({"ID_test_roc_auc": agg_func, "OOD_test_roc_auc": agg_func}).reset_index()
# Flatten column names
df_stats.columns = ["model", f"ID_{agg_func}", f"OOD_{agg_func}"]
# Calculate gap between ID and OOD
df_stats["gap"] = df_stats[f"ID_{agg_func}"] - df_stats[f"OOD_{agg_func}"]
# Sort based on MODEL_MAPPING order
df_stats = sort_df_by_model_mapping(df_stats, MODEL_MAPPING)
# Set up bar positions
x = np.arange(len(df_stats))
width = 0.35
# Create bars
axes[i].bar(x - width / 2, df_stats[f"ID_{agg_func}"], width, label="ID Test")
axes[i].bar(x + width / 2, df_stats[f"OOD_{agg_func}"], width, label="OOD Test")
# Add gap values as text above bars
for j, gap in enumerate(df_stats["gap"]):
# Position text above the higher of the two bars
y_pos = max(df_stats[f"ID_{agg_func}"].iloc[j], df_stats[f"OOD_{agg_func}"].iloc[j])
gap_text = f"{gap:.3f}"
axes[i].text(x[j], y_pos + 0.01, gap_text, ha="center", va="bottom", fontsize=10)
# Customize plot
axes[i].set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=22, pad=10)
axes[i].set_xticks(x)
axes[i].set_xticklabels(
[r"\textbf{" + MODEL_MAPPING[model] + "}" for model in df_stats["model"]], rotation=45, ha="right", fontsize=14
)
axes[i].set_ylabel(r"\textbf{ROC AUC}", fontsize=16, labelpad=10)
axes[i].grid(True, axis="y", linestyle="--", alpha=0.7) # Only show y-axis grid
axes[i].grid(False, axis="x") # Explicitly turn off x-axis grid
axes[i].set_ylim(0.5, 0.9) # Adjust as needed
# get legend, put it to the bottom right of the figure
handles, labels = axes[0].get_legend_handles_labels()
legend = fig.legend(
handles=handles,
labels=labels,
loc="lower right",
bbox_to_anchor=(0.8, 0.13),
title=r"\textbf{Test set type}",
fontsize=20,
title_fontsize=22,
)
# Remove any unused subplots
if len(axes) > len(splitters):
for j in range(len(splitters), len(axes)):
axes[j].remove()
plt.suptitle(r"$\textbf{Model Performance Across Different Splitting Methods}$", fontsize=26, y=1.05)
if save:
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.pdf"), bbox_inches="tight"
)
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.png"),
bbox_inches="tight",
dpi=600,
)
plt.show()
In [ ]:
Copied!
df_split = benchmark_results[benchmark_results["split"] == splitter]
df_split
df_split = benchmark_results[benchmark_results["split"] == splitter]
df_split
In [ ]:
Copied!
df_stats
df_stats
In [ ]:
Copied!
# Create plot with subplots for each splitter
save = False
sns.set_palette("Set2")
agg_func = "mean"
splitter = "random"
# Set up the figure and grid
fig, ax = plt.subplots(figsize=(8, 6))
# Filter data for this splitter
df_split = benchmark_results[benchmark_results["split"] == splitter]
df_split = df_split[df_split["model"] == "gem"]
# Calculate mean and std values for each model
df_stats = (
df_split.groupby("dataset")
.agg({"ID_test_roc_auc": [agg_func, "std"], "OOD_test_roc_auc": [agg_func, "std"]})
.reset_index()
)
# Flatten column names
df_stats.columns = ["dataset", f"ID_{agg_func}", "ID_std", f"OOD_{agg_func}", "OOD_std"]
# Calculate gap between ID and OOD
df_stats["gap"] = df_stats[f"ID_{agg_func}"] - df_stats[f"OOD_{agg_func}"]
# Set up bar positions
x = np.arange(len(df_stats))
width = 0.35
# Create bars with error bars
ax.bar(
x - width / 2,
df_stats[f"ID_{agg_func}"],
width,
label="ID Test",
yerr=df_stats["ID_std"],
capsize=5,
error_kw=dict(capthick=2),
)
ax.bar(
x + width / 2,
df_stats[f"OOD_{agg_func}"],
width,
label="OOD Test",
yerr=df_stats["OOD_std"],
capsize=5,
error_kw=dict(capthick=2),
)
# Add gap values as text above bars
for j, gap in enumerate(df_stats["gap"]):
# Position text above the higher of the two bars plus its error bar
y_pos = max(
df_stats[f"ID_{agg_func}"].iloc[j] + df_stats["ID_std"].iloc[j],
df_stats[f"OOD_{agg_func}"].iloc[j] + df_stats["OOD_std"].iloc[j],
)
gap_text = f"{gap:.3f}"
ax.text(x[j], y_pos + 0.01, gap_text, ha="center", va="bottom", fontsize=10)
# Customize plot
ax.set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=22, pad=10)
ax.set_xticks(x)
ax.set_xticklabels(
[r"\textbf{" + dataset + "}" for dataset in df_stats["dataset"]], rotation=45, ha="right", fontsize=14
)
ax.set_ylabel(r"\textbf{ROC AUC}", fontsize=16, labelpad=10)
ax.grid(True, axis="y", linestyle="--", alpha=0.7) # Only show y-axis grid
ax.grid(False, axis="x") # Explicitly turn off x-axis grid
ax.set_ylim(0.6, 1) # Adjust as needed
# get legend, put it to the bottom right of the figure
handles, labels = axes[0].get_legend_handles_labels()
legend = fig.legend(
handles=handles,
labels=labels,
loc="lower right",
bbox_to_anchor=(0.8, 0.13),
title=r"\textbf{Test set type}",
fontsize=20,
title_fontsize=22,
)
plt.suptitle(r"$\textbf{Model Performance Across Different Splitting Methods}$", fontsize=26, y=1.05)
if save:
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.pdf"), bbox_inches="tight"
)
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.png"),
bbox_inches="tight",
dpi=600,
)
plt.show()
# Create plot with subplots for each splitter
save = False
sns.set_palette("Set2")
agg_func = "mean"
splitter = "random"
# Set up the figure and grid
fig, ax = plt.subplots(figsize=(8, 6))
# Filter data for this splitter
df_split = benchmark_results[benchmark_results["split"] == splitter]
df_split = df_split[df_split["model"] == "gem"]
# Calculate mean and std values for each model
df_stats = (
df_split.groupby("dataset")
.agg({"ID_test_roc_auc": [agg_func, "std"], "OOD_test_roc_auc": [agg_func, "std"]})
.reset_index()
)
# Flatten column names
df_stats.columns = ["dataset", f"ID_{agg_func}", "ID_std", f"OOD_{agg_func}", "OOD_std"]
# Calculate gap between ID and OOD
df_stats["gap"] = df_stats[f"ID_{agg_func}"] - df_stats[f"OOD_{agg_func}"]
# Set up bar positions
x = np.arange(len(df_stats))
width = 0.35
# Create bars with error bars
ax.bar(
x - width / 2,
df_stats[f"ID_{agg_func}"],
width,
label="ID Test",
yerr=df_stats["ID_std"],
capsize=5,
error_kw=dict(capthick=2),
)
ax.bar(
x + width / 2,
df_stats[f"OOD_{agg_func}"],
width,
label="OOD Test",
yerr=df_stats["OOD_std"],
capsize=5,
error_kw=dict(capthick=2),
)
# Add gap values as text above bars
for j, gap in enumerate(df_stats["gap"]):
# Position text above the higher of the two bars plus its error bar
y_pos = max(
df_stats[f"ID_{agg_func}"].iloc[j] + df_stats["ID_std"].iloc[j],
df_stats[f"OOD_{agg_func}"].iloc[j] + df_stats["OOD_std"].iloc[j],
)
gap_text = f"{gap:.3f}"
ax.text(x[j], y_pos + 0.01, gap_text, ha="center", va="bottom", fontsize=10)
# Customize plot
ax.set_title(rf"$\textbf{{{SPLIT_TYPPE_MAPPING[splitter]}}}$", fontsize=22, pad=10)
ax.set_xticks(x)
ax.set_xticklabels(
[r"\textbf{" + dataset + "}" for dataset in df_stats["dataset"]], rotation=45, ha="right", fontsize=14
)
ax.set_ylabel(r"\textbf{ROC AUC}", fontsize=16, labelpad=10)
ax.grid(True, axis="y", linestyle="--", alpha=0.7) # Only show y-axis grid
ax.grid(False, axis="x") # Explicitly turn off x-axis grid
ax.set_ylim(0.6, 1) # Adjust as needed
# get legend, put it to the bottom right of the figure
handles, labels = axes[0].get_legend_handles_labels()
legend = fig.legend(
handles=handles,
labels=labels,
loc="lower right",
bbox_to_anchor=(0.8, 0.13),
title=r"\textbf{Test set type}",
fontsize=20,
title_fontsize=22,
)
plt.suptitle(r"$\textbf{Model Performance Across Different Splitting Methods}$", fontsize=26, y=1.05)
if save:
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.pdf"), bbox_inches="tight"
)
plt.savefig(
os.path.join(CHECKOUT_PATH, "assets", "figures", "model_performance_by_splitter.png"),
bbox_inches="tight",
dpi=600,
)
plt.show()
Supporting Information¶
Datasets PhysicoChemical Properties (Figure S3)¶
In [ ]:
Copied!
# Function to compute physicochemical properties
def compute_properties(smiles):
mol = Chem.MolFromSmiles(smiles)
if mol:
return {
"Molecular_Weight": Descriptors.MolWt(mol),
"LogP": Descriptors.MolLogP(mol),
"TPSA": Descriptors.TPSA(mol),
"HBD": Descriptors.NumHDonors(mol),
"HBA": Descriptors.NumHAcceptors(mol),
"Rotatable_Bonds": Descriptors.NumRotatableBonds(mol),
}
else:
return {
"Molecular_Weight": None,
"LogP": None,
"TPSA": None,
"HBD": None,
"HBA": None,
"Rotatable_Bonds": None,
}
# Function to process a dataset
def process_dataset(df):
properties_df = df["smiles"].apply(compute_properties).apply(pd.Series)
return pd.concat([df, properties_df], axis=1)
# Function to compute physicochemical properties
def compute_properties(smiles):
mol = Chem.MolFromSmiles(smiles)
if mol:
return {
"Molecular_Weight": Descriptors.MolWt(mol),
"LogP": Descriptors.MolLogP(mol),
"TPSA": Descriptors.TPSA(mol),
"HBD": Descriptors.NumHDonors(mol),
"HBA": Descriptors.NumHAcceptors(mol),
"Rotatable_Bonds": Descriptors.NumRotatableBonds(mol),
}
else:
return {
"Molecular_Weight": None,
"LogP": None,
"TPSA": None,
"HBD": None,
"HBA": None,
"Rotatable_Bonds": None,
}
# Function to process a dataset
def process_dataset(df):
properties_df = df["smiles"].apply(compute_properties).apply(pd.Series)
return pd.concat([df, properties_df], axis=1)
In [ ]:
Copied!
props_df = []
for datasets in DATASET_NAMES:
df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", datasets, f"{datasets}_standardize.csv"))
processed_df = process_dataset(df)
props_df.append(processed_df)
datasets = {}
for i, dataset in enumerate(DATASET_NAMES):
datasets[dataset] = props_df[i]
props_df = []
for datasets in DATASET_NAMES:
df = pd.read_csv(os.path.join(DATASET_PATH, "TDC", datasets, f"{datasets}_standardize.csv"))
processed_df = process_dataset(df)
props_df.append(processed_df)
datasets = {}
for i, dataset in enumerate(DATASET_NAMES):
datasets[dataset] = props_df[i]
In [ ]:
Copied!
from alinemol.utils.plot_utils import plot_radar_subplots
plot_radar_subplots(datasets, save=False)
from alinemol.utils.plot_utils import plot_radar_subplots
plot_radar_subplots(datasets, save=False)
Datasets Size Comparisson For Differnt Splitters (ID vs OOD test size) (Figure S4)¶
In [ ]:
Copied!
# Fixes split type. Then, find out ratio of test set size to the size of the dataset
dataset_category = "TDC"
dataset_names = DATASET_NAMES
# dataset_names = "CYP1A2"
split_types = SPLIT_TYPES
dfs = []
for dataset_name in dataset_names:
for split_type in split_types:
dataset_folder = os.path.join(DATASET_PATH, dataset_category, dataset_name, "split", split_type)
id_size = []
id_test_size = []
ood_test_size = []
df = pd.DataFrame()
with open(os.path.join(dataset_folder, "config.json"), "r") as f:
data_config = json.load(f)
for i in range(10):
id_size.append(data_config[f"train_size_{i}"])
ood_test_size.append(data_config[f"test_size_{i}"])
id_test_size.append(data_config[f"train_size_{i}"] * 0.2)
id_frac = [id_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
ood_test_frac = [ood_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
id_test_frac = [id_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
df["ood_test_size"] = ood_test_frac
df["id_test_size"] = id_test_frac
df["split_type"] = [split_type] * 10
df["dataset"] = [dataset_name] * 10
dfs.append(df)
df = pd.concat(dfs)
df.head()
# Fixes split type. Then, find out ratio of test set size to the size of the dataset
dataset_category = "TDC"
dataset_names = DATASET_NAMES
# dataset_names = "CYP1A2"
split_types = SPLIT_TYPES
dfs = []
for dataset_name in dataset_names:
for split_type in split_types:
dataset_folder = os.path.join(DATASET_PATH, dataset_category, dataset_name, "split", split_type)
id_size = []
id_test_size = []
ood_test_size = []
df = pd.DataFrame()
with open(os.path.join(dataset_folder, "config.json"), "r") as f:
data_config = json.load(f)
for i in range(10):
id_size.append(data_config[f"train_size_{i}"])
ood_test_size.append(data_config[f"test_size_{i}"])
id_test_size.append(data_config[f"train_size_{i}"] * 0.2)
id_frac = [id_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
ood_test_frac = [ood_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
id_test_frac = [id_test_size[i] / (id_size[i] + ood_test_size[i]) * 100 for i in range(10)]
df["ood_test_size"] = ood_test_frac
df["id_test_size"] = id_test_frac
df["split_type"] = [split_type] * 10
df["dataset"] = [dataset_name] * 10
dfs.append(df)
df = pd.concat(dfs)
df.head()
In [ ]:
Copied!
from alinemol.utils.plot_utils import boxplot_test_set_size
boxplot_test_set_size(df, save=False)
from alinemol.utils.plot_utils import boxplot_test_set_size
boxplot_test_set_size(df, save=False)
Figure S7-10¶
In [ ]:
Copied!
# Figure S7, S8
from alinemol.utils.plot_utils import heatmap_plot_id_ood
heatmap_plot_id_ood(results=benchmark_results, metric="roc_auc", save=True) # Figure S7
heatmap_plot_id_ood(results=benchmark_results, metric="accuracy", save=True) # Figure S8
# Figure S7, S8
from alinemol.utils.plot_utils import heatmap_plot_id_ood
heatmap_plot_id_ood(results=benchmark_results, metric="roc_auc", save=True) # Figure S7
heatmap_plot_id_ood(results=benchmark_results, metric="accuracy", save=True) # Figure S8
In [ ]:
Copied!
# Figure S9
from alinemol.utils.plot_utils import heatmap_plot
heatmap_plot(results=benchmark_results, metric="accuracy", save=True)
# Figure S9
from alinemol.utils.plot_utils import heatmap_plot
heatmap_plot(results=benchmark_results, metric="accuracy", save=True)
In [ ]:
Copied!
# Figure S10
from alinemol.utils.plot_utils import heatmap_plot_all_dataset
heatmap_plot_all_dataset(results=benchmark_results, metric="accuracy", save=True)
# Figure S10
from alinemol.utils.plot_utils import heatmap_plot_all_dataset
heatmap_plot_all_dataset(results=benchmark_results, metric="accuracy", save=True)
Activity Ratio Plot (Figure S11)¶
In [ ]:
Copied!
# Load and plot activity ratios for each dataset split
save = True
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
import glob
import os
# Set style for a professional look
sns.set_palette("Set2")
# Create figure and subplots
fig = plt.figure(figsize=(24, 16))
gs = fig.add_gridspec(3, 3, hspace=1.1, wspace=0.2)
# Create list to store all plot data
all_plot_data = []
# List of datasets
datasets = DATASET_NAMES
for idx, dataset in enumerate(datasets):
dataset_plot_data = []
for split in SPLIT_TYPES:
# Get train and test files
train_files = sorted(glob.glob(f"datasets/TDC/{dataset}/split/{split}/train_*.csv"))
test_files = sorted(glob.glob(f"datasets/TDC/{dataset}/split/{split}/test_*.csv"))
# Calculate activity ratios
for train_file, test_file in zip(train_files, test_files):
train_df = pd.read_csv(train_file)
test_df = pd.read_csv(test_file)
train_ratio = (train_df["label"] == 1).mean() * 100
test_ratio = (test_df["label"] == 1).mean() * 100
# Store data
dataset_plot_data.extend(
[
{"Dataset": dataset, "Split": split, "Set": "Train", "Activity Ratio (%)": train_ratio},
{"Dataset": dataset, "Split": split, "Set": "Test", "Activity Ratio (%)": test_ratio},
]
)
all_plot_data.extend(dataset_plot_data)
# Create subplot
ax = fig.add_subplot(gs[idx // 3, idx % 3])
# Create plot data
plot_df = pd.DataFrame(dataset_plot_data)
# Create boxplot with enhanced styling
sns.boxplot(
data=plot_df,
x="Split",
y="Activity Ratio (%)",
hue="Set",
ax=ax,
boxprops={"alpha": 0.8},
showfliers=False,
medianprops={"linewidth": 1.5},
width=0.8,
capwidths=0.25,
) # Hide outliers for cleaner look
# Customize subplot
ax.set_title(rf"\textbf{{{dataset}}}", fontsize=20, pad=12)
if idx in [5, 6, 7]:
ax.set_xlabel(r"$\textbf{Split}$", fontsize=20, labelpad=12)
else:
ax.set_xlabel("")
if idx in [0, 3, 6]: # Only show y-label for leftmost plots
ax.set_ylabel(r"$\textbf{Activity Ratio (\%)}$", fontsize=20, labelpad=12)
else:
ax.set_ylabel("")
ax.grid(True, alpha=0.3, axis="both")
# format ticklabels
ax.set_xticks(range(len(ax.get_xticklabels())))
ax.set_xticklabels(
[rf"\textbf{{{SPLIT_TYPPE_MAPPING[x.get_text()]}}}" for x in ax.get_xticklabels()],
fontsize=16,
ha="right",
rotation=45,
)
plt.setp(ax.get_yticklabels(), fontsize=16)
# set ylim 10-70
ax.set_ylim(10, 70)
# remove individual legend
ax.get_legend().remove()
# set shared legend with custom labels
handles, labels = ax.get_legend_handles_labels()
labels = [r"\textbf{Training Set}", r"\textbf{OOD Test Set}"] # Custom legend labels
fig.legend(
handles, labels, title=r"\textbf{Set}", frameon=True, bbox_to_anchor=(0.74, 0.22), title_fontsize=24, fontsize=18
)
# Add overall title
fig.suptitle(r"\textbf{Activity Ratios Across Datasets and Splits}", fontsize=24, y=0.95) # Removed \textbf
# Adjust layout
plt.tight_layout()
if save:
plt.savefig("assets/figures/activity_ratios.png", dpi=300, bbox_inches="tight")
plt.show()
# Load and plot activity ratios for each dataset split
save = True
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
import glob
import os
# Set style for a professional look
sns.set_palette("Set2")
# Create figure and subplots
fig = plt.figure(figsize=(24, 16))
gs = fig.add_gridspec(3, 3, hspace=1.1, wspace=0.2)
# Create list to store all plot data
all_plot_data = []
# List of datasets
datasets = DATASET_NAMES
for idx, dataset in enumerate(datasets):
dataset_plot_data = []
for split in SPLIT_TYPES:
# Get train and test files
train_files = sorted(glob.glob(f"datasets/TDC/{dataset}/split/{split}/train_*.csv"))
test_files = sorted(glob.glob(f"datasets/TDC/{dataset}/split/{split}/test_*.csv"))
# Calculate activity ratios
for train_file, test_file in zip(train_files, test_files):
train_df = pd.read_csv(train_file)
test_df = pd.read_csv(test_file)
train_ratio = (train_df["label"] == 1).mean() * 100
test_ratio = (test_df["label"] == 1).mean() * 100
# Store data
dataset_plot_data.extend(
[
{"Dataset": dataset, "Split": split, "Set": "Train", "Activity Ratio (%)": train_ratio},
{"Dataset": dataset, "Split": split, "Set": "Test", "Activity Ratio (%)": test_ratio},
]
)
all_plot_data.extend(dataset_plot_data)
# Create subplot
ax = fig.add_subplot(gs[idx // 3, idx % 3])
# Create plot data
plot_df = pd.DataFrame(dataset_plot_data)
# Create boxplot with enhanced styling
sns.boxplot(
data=plot_df,
x="Split",
y="Activity Ratio (%)",
hue="Set",
ax=ax,
boxprops={"alpha": 0.8},
showfliers=False,
medianprops={"linewidth": 1.5},
width=0.8,
capwidths=0.25,
) # Hide outliers for cleaner look
# Customize subplot
ax.set_title(rf"\textbf{{{dataset}}}", fontsize=20, pad=12)
if idx in [5, 6, 7]:
ax.set_xlabel(r"$\textbf{Split}$", fontsize=20, labelpad=12)
else:
ax.set_xlabel("")
if idx in [0, 3, 6]: # Only show y-label for leftmost plots
ax.set_ylabel(r"$\textbf{Activity Ratio (\%)}$", fontsize=20, labelpad=12)
else:
ax.set_ylabel("")
ax.grid(True, alpha=0.3, axis="both")
# format ticklabels
ax.set_xticks(range(len(ax.get_xticklabels())))
ax.set_xticklabels(
[rf"\textbf{{{SPLIT_TYPPE_MAPPING[x.get_text()]}}}" for x in ax.get_xticklabels()],
fontsize=16,
ha="right",
rotation=45,
)
plt.setp(ax.get_yticklabels(), fontsize=16)
# set ylim 10-70
ax.set_ylim(10, 70)
# remove individual legend
ax.get_legend().remove()
# set shared legend with custom labels
handles, labels = ax.get_legend_handles_labels()
labels = [r"\textbf{Training Set}", r"\textbf{OOD Test Set}"] # Custom legend labels
fig.legend(
handles, labels, title=r"\textbf{Set}", frameon=True, bbox_to_anchor=(0.74, 0.22), title_fontsize=24, fontsize=18
)
# Add overall title
fig.suptitle(r"\textbf{Activity Ratios Across Datasets and Splits}", fontsize=24, y=0.95) # Removed \textbf
# Adjust layout
plt.tight_layout()
if save:
plt.savefig("assets/figures/activity_ratios.png", dpi=300, bbox_inches="tight")
plt.show()
In [ ]:
Copied!
In [ ]:
Copied!