Interpretable deep survival analysis of Alzheimer's disease via metabolic genetic variants.
The 5 matches
- [1] § 3 Results › 3.1 5-Fold cross-validation and performance evaluation ↔ visualize.py, lines 600–665 · score 0.61 · Feedforward Neural Network, Fold Cross Validation, Linear Regression, Standard Deviation, Std, LR
- [2] § 3 Results › 3.1 5-Fold cross-validation and performance evaluation ↔ visualize.py, lines 600–665 · score 0.56 · fold cross validation, linear regression, standard deviation, training, FFN, models
- [3] § 2 Methods › 2.1 Data collection and preprocessing ↔ visualize.py, lines 491–557 · score 0.54 · T2D, diabetes, dyslipidemia, chromosome, position, Variant
- [4] § 3 Results › 3.2 Feature importance analysis using SHAP ↔ visualize.py, lines 691–762 · score 0.52 · APOE E3, APOE E2, APOE E4, MMSE, model
- [5] § 2 Methods › 2.3 Evaluating and comparing model with baseline models ↔ visualize.py, lines 45–57 · score 0.51 · Weibull AFT, linear model, concordance, score, event
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 888 lines · 28 KB · MIT · 5 matches
- # %%
- import os
- from dataclasses import dataclass
- from itertools import combinations
- from pathlib import Path
- import matplotlib.pyplot as plt
- import numpy as np
- import pandas as pd
- import seaborn as sns
- import torch
- from captum.attr import IntegratedGradients
- from ignite.handlers import Checkpoint
- from lifelines import KaplanMeierFitter, WeibullAFTFitter
- from model import DeepWeibullModel
- from scipy.stats import chi2_contingency, fisher_exact
- from seaborn import catplot, boxplot
- from sklearn.cluster import (
- BisectingKMeans,
- )
- from sklearn.decomposition import PCA
- from torchsurv.loss.weibull import log_hazard, survival_function
- from torchsurv.metrics.cindex import ConcordanceIndex
- from torchsurv.metrics.brier_score import BrierScore
- from torchsurv.stats.kaplan_meier import KaplanMeierEstimator
- from lifelines.statistics import logrank_test
- from shap import DeepExplainer
- import shap
- MMSE_COLUMN_NAME = "MMSE"
- TARGET_COLUMN_NAME = "Diagnose"
- AGE_COLUMN_NAME = "Age"
- SEX_COLUMN_NAME = "Sex"
- APOE_COLUMN_NAME = "APOE"
- FOLD_COLUMN_NAME = "Fold"
- CLUSTER_COLUMN_NAME = "Cluster"
- BASE_AGE = 55
- EPOCHS = 10_000
- plt.rcParams["font.family"] = "DejaVu Serif"
- plt.rcParams["font.serif"] = ["Book"]
- plt.rcParams["mathtext.fontset"] = "cm"
- def get_linear_aft_result(df_train, df_test) -> tuple[pd.DataFrame, float, float]:
- weibull_aft = WeibullAFTFitter()
- weibull_aft.fit(
- df_train, duration_col=AGE_COLUMN_NAME, event_col=TARGET_COLUMN_NAME
- )
- linear_model_c_index_test = weibull_aft.score(
- df_test, scoring_method="concordance_index"
- )
- return (
- weibull_aft.summary,
- weibull_aft.concordance_index_,
- linear_model_c_index_test,
- )
- def get_fnn_result(df_train, df_test, deep_weibull_model):
- features_train_df = df_train.drop(columns=[TARGET_COLUMN_NAME, AGE_COLUMN_NAME])
- features_train = torch.tensor(features_train_df.values, dtype=torch.float32)
- log_scale_train = deep_weibull_model.forward(features_train)
- log_shape_train = deep_weibull_model.log_shape().repeat(log_scale_train.size(0), 1)
- log_params_train = torch.cat((log_scale_train, log_shape_train), dim=-1)
- durations_train = torch.tensor(
- df_train[AGE_COLUMN_NAME].values, dtype=torch.float32
- )
- events_train = torch.tensor(df_train[TARGET_COLUMN_NAME].values, dtype=torch.bool)
- log_hz = log_hazard(log_params_train, durations_train)
- c_index = ConcordanceIndex()
- c_index_train = c_index(log_hz, events_train, durations_train)
- surv_train = survival_function(log_params_train, durations_train)
- brier_score = BrierScore()
- brier_train = brier_score(surv_train, events_train, durations_train)
- brier_score_train = brier_score.integral()
- log_scale_test = deep_weibull_model.forward(
- torch.tensor(
- df_test.drop(columns=[TARGET_COLUMN_NAME, AGE_COLUMN_NAME]).values,
- dtype=torch.float32,
- )
- )
- log_shape_test = deep_weibull_model.log_shape().repeat(log_scale_test.size(0), 1)
- log_params_test = torch.cat((log_scale_test, log_shape_test), dim=-1)
- durations_test = torch.tensor(df_test[AGE_COLUMN_NAME].values, dtype=torch.float32)
- events_test = torch.tensor(df_test[TARGET_COLUMN_NAME].values, dtype=torch.bool)
- log_hz_test = log_hazard(log_params_test, durations_test)
- c_index_test = c_index(log_hz_test, events_test, durations_test)
- surv_test = survival_function(log_params_test, durations_test)
- brier_score_test = brier_score(surv_test, events_test, durations_test)
- brier_score_test = brier_score.integral()
- return c_index_train, c_index_test, brier_score_train, brier_score_test
- def get_IG(model: DeepWeibullModel, features: pd.DataFrame) -> pd.DataFrame:
- features = features.drop(
- columns=[TARGET_COLUMN_NAME, AGE_COLUMN_NAME, FOLD_COLUMN_NAME]
- )
- column_names = features.columns.tolist()
- features = torch.tensor(features.values, dtype=torch.float32)
- ig = IntegratedGradients(model)
- base_line = torch.zeros((1, features.shape[1]), dtype=torch.float32)
- # base_line = features.mean(dim=0, keepdim=True)
- attributions = ig.attribute(
- features, baselines=base_line.repeat([features.shape[0], 1]), n_steps=50
- )
- attributions_df = pd.DataFrame(
- attributions.cpu().detach().numpy(), columns=column_names
- )
- return attributions_df
- def get_shap(
- model: DeepWeibullModel, features: pd.DataFrame, seed: int, background_size=900
- ) -> pd.DataFrame:
- column_names = features.columns.tolist()
- features = torch.tensor(features.values, dtype=torch.float32)
- rng = np.random.default_rng(seed)
- background_samples = rng.choice(
- features.shape[0], size=background_size, replace=False
- )
- explainer = DeepExplainer(
- model, features[background_samples]
- ) # Use a subset as background
- shap_values = explainer.shap_values(features)
- shap_values = np.array(shap_values.squeeze(-1))
- attributions_df = pd.DataFrame(shap_values, columns=column_names)
- plt.figure(figsize=(10, 6))
- shap.summary_plot(
- shap_values,
- features.cpu().numpy(),
- feature_names=column_names,
- )
- plt.xlim(-0.5, 0.5)
- plt.tight_layout()
- plt.savefig("fig_img/shap_beeswarm.pdf", dpi=300, bbox_inches="tight")
- return attributions_df
- def get_shap_specific_interaction(
- model,
- features: pd.DataFrame,
- feature_x: str,
- feature_color: str,
- seed: int,
- background_size=900,
- ) -> pd.DataFrame:
- column_names = features.columns.tolist()
- features_tensor = torch.tensor(features.values, dtype=torch.float32)
- rng = np.random.default_rng(seed)
- background_samples = rng.choice(
- features_tensor.shape[0], size=background_size, replace=False
- )
- explainer = shap.DeepExplainer(model, features_tensor[background_samples])
- shap_values = explainer.shap_values(features_tensor)
- if isinstance(shap_values, list):
- shap_values = shap_values[0]
- if len(shap_values.shape) > 2:
- shap_values = np.array(shap_values.squeeze(-1))
- attributions_df = pd.DataFrame(shap_values, columns=column_names)
- plt.figure(figsize=(8, 5))
- shap.dependence_plot(
- feature_x,
- shap_values,
- features_tensor.cpu().numpy(),
- feature_names=column_names,
- interaction_index=feature_color,
- show=False,
- x_jitter=0.1,
- alpha=0.5,
- )
- plt.tight_layout()
- plt.savefig(
- f"fig_img/shap_dependence_{feature_x}_vs_{feature_color}.pdf",
- dpi=300,
- bbox_inches="tight",
- )
- plt.close()
- return attributions_df
- def get_mult_shap_specific_interaction(
- model,
- features: pd.DataFrame,
- feature_xs: list[str],
- feature_colors: list[str],
- seed: int,
- background_size=900,
- ) -> pd.DataFrame:
- column_names = features.columns.tolist()
- features_tensor = torch.tensor(features.values, dtype=torch.float32)
- rng = np.random.default_rng(seed)
- background_samples = rng.choice(
- features_tensor.shape[0], size=background_size, replace=False
- )
- explainer = shap.DeepExplainer(model, features_tensor[background_samples])
- shap_values = explainer.shap_values(features_tensor)
- if isinstance(shap_values, list):
- shap_values = shap_values[0]
- if len(shap_values.shape) > 2:
- shap_values = np.array(shap_values.squeeze(-1))
- attributions_df = pd.DataFrame(shap_values, columns=column_names)
- fig, axes = plt.subplots(
- len(feature_xs),
- len(feature_colors),
- figsize=(5 * len(feature_colors), 4 * len(feature_xs)),
- )
- for i, feature_x in enumerate(feature_xs):
- for j, feature_color in enumerate(feature_colors):
- shap.dependence_plot(
- feature_x,
- shap_values,
- features_tensor.cpu().numpy(),
- feature_names=column_names,
- interaction_index=feature_color,
- show=False,
- x_jitter=0.1,
- alpha=0.5,
- ax=axes[i, j],
- )
- plt.tight_layout()
- plt.savefig(
- f"fig_img/shap_inter.pdf",
- dpi=300,
- bbox_inches="tight",
- )
- plt.close()
- return attributions_df
- def cluster_attributions(features_ig: pd.DataFrame, n_clusters=5, seed=42):
- clusterer = BisectingKMeans(
- n_clusters=n_clusters,
- random_state=seed,
- n_init=1,
- max_iter=10_000_000,
- )
- clusterer.fit(features_ig)
- features_ig[CLUSTER_COLUMN_NAME] = clusterer.labels_
- return features_ig
- def plot_pca(
- ax,
- attributions_df: pd.DataFrame,
- palette: dict,
- x_lim=[-1.25, 1.25],
- y_lim=[-0.25, 0.25],
- ):
- pca = PCA(n_components=2)
- attributions_pca = pca.fit_transform(attributions_df.drop(columns="Cluster"))
- scatter_ax = sns.scatterplot(
- x=attributions_pca[:, 0],
- y=attributions_pca[:, 1],
- hue=attributions_df["Cluster"],
- palette=palette,
- alpha=0.6,
- s=60,
- ax=ax,
- )
- scatter_ax.set_xlim(x_lim)
- scatter_ax.set_ylim(y_lim)
- handles, labels = scatter_ax.get_legend_handles_labels()
- order = sorted(range(len(labels)), key=lambda k: labels[k])
- scatter_ax.legend(
- [handles[idx] for idx in order],
- [labels[idx] for idx in order],
- title="Cluster",
- loc="upper left",
- )
- return scatter_ax
- def median_time_km(clusters, durations, events):
- aucs = {}
- for cluster in clusters.unique():
- mask = clusters == cluster
- if events[mask].sum() == 0:
- aucs[cluster] = 1.0 * (durations[mask].max() - BASE_AGE)
- continue
- km = KaplanMeierEstimator()
- km_fitter = KaplanMeierFitter()
- km_fitter.fit(
- durations[mask],
- event_observed=events[mask],
- )
- aucs[cluster] = km_fitter.median_survival_time_
- sorted_aucs = dict(sorted(aucs.items(), key=lambda item: item[1], reverse=True))
- return sorted_aucs
- def plot_km(
- ax,
- clusters: pd.Series | None,
- durations: np.ndarray,
- events: np.ndarray,
- colors: list,
- cluster_list=None,
- total=False,
- total_color="black",
- ):
- if clusters is None:
- clusters = pd.Series(["All"] * len(durations))
- for cluster, color in zip(cluster_list or clusters.unique(), colors):
- mask = clusters == cluster
- count = mask.sum()
- if events[mask].sum() == 0:
- ax.plot(
- durations[mask] + BASE_AGE,
- np.ones_like(durations[mask]),
- linestyle=":",
- label=f"{cluster} (n={count}) - No events",
- color=color,
- )
- else:
- km = KaplanMeierEstimator()
- km(
- torch.tensor(events[mask], dtype=torch.bool),
- torch.tensor(durations[mask], dtype=torch.float32),
- )
- times = (
- km.time.cpu().numpy() + BASE_AGE
- if hasattr(km.time, "cpu")
- else km.time.numpy() + BASE_AGE
- )
- surv = (
- km.km_est.cpu().numpy()
- if hasattr(km.km_est, "cpu")
- else km.km_est.numpy()
- )
- ax.step(
- times, surv, where="post", label=f"{cluster} (n={count})", color=color
- )
- if cluster_list and len(cluster_list) == 2:
- mask0 = clusters == cluster_list[0]
- mask1 = clusters == cluster_list[1]
- logrank_test_result = logrank_test(
- df_train[AGE_COLUMN_NAME][mask0],
- df_train[AGE_COLUMN_NAME][mask1],
- event_observed_A=df_train[TARGET_COLUMN_NAME][mask0],
- event_observed_B=df_train[TARGET_COLUMN_NAME][mask1],
- )
- p_value = logrank_test_result.p_value
- ax.plot([], [], " ", label=f"p-value = {p_value:.3g}")
- if total:
- km = KaplanMeierEstimator()
- km(
- torch.tensor(events, dtype=torch.bool),
- torch.tensor(durations, dtype=torch.float32),
- )
- times = (
- km.time.cpu().numpy() + BASE_AGE
- if hasattr(km.time, "cpu")
- else km.time.numpy() + BASE_AGE
- )
- surv = (
- km.km_est.cpu().numpy() if hasattr(km.km_est, "cpu") else km.km_est.numpy()
- )
- count = len(durations)
- ax.step(
- times, surv, where="post", label=f"Total (n={count})", color=total_color
- )
- if clusters.nunique() > 0:
- handles, labels = ax.get_legend_handles_labels()
- order = sorted(range(len(labels)), key=lambda k: labels[k])
- if len(labels) > 1:
- ax.legend(
- [handles[idx] for idx in order],
- [labels[idx] for idx in order],
- title="Cluster",
- loc="lower left",
- )
- else:
- legend = ax.get_legend()
- if legend:
- legend.remove()
- return ax
- def make_crosstab(df: pd.DataFrame) -> pd.DataFrame:
- df_long = df.melt(
- id_vars=["Cluster"],
- value_vars=df.drop(columns=["Cluster"]).columns.to_list(),
- var_name="feature",
- value_name="value",
- )
- counts_df = pd.crosstab(
- index=[df_long["Cluster"], df_long["feature"]], columns=df_long["value"]
- )
- total_counts_df = pd.crosstab(index=[df_long["feature"]], columns=df_long["value"])
- total_counts_df.columns = ["0", "1"]
- counts_df.columns = ["0", "1"]
- counts_df = counts_df[["1", "0"]]
- total_counts_df = total_counts_df[["1", "0"]]
- comparison_clusters = list(
- combinations(["Total"] + df["Cluster"].unique().tolist(), 2)
- )
- results = []
- for feature in df.drop(columns=["Cluster"]).columns:
- for cluster2 in df["Cluster"].unique().tolist():
- cluster1 = "Total"
- try:
- if cluster1 == "Total":
- base_counts = total_counts_df.loc[(feature)]
- else:
- base_counts = counts_df.loc[(cluster1, feature)]
- if cluster2 == "Total":
- comp_counts = total_counts_df.loc[(feature)]
- else:
- comp_counts = counts_df.loc[(cluster2, feature)]
- contingency_table = pd.DataFrame([base_counts, comp_counts])
- contingency_table.index = [
- f"Cluster {cluster1}",
- f"Cluster {cluster2}",
- ]
- chi2, p_value, dof, expected = chi2_contingency(contingency_table)
- mat = contingency_table.values.astype(float)
- odds_ratio = (
- (mat[0, 1] * mat[1, 0]) / (mat[0, 0] * mat[1, 1])
- if mat[0, 0] * mat[1, 1] != 0.0
- else np.nan
- )
- if odds_ratio is np.nan or odds_ratio == 0:
- mat = contingency_table.values.astype(float) + 0.5
- odds_ratio = (mat[0, 1] * mat[1, 0]) / (mat[0, 0] * mat[1, 1])
- results.append(
- {
- "feature": feature,
- "comparison": f"{cluster1} vs {cluster2}",
- "chi2_statistic": chi2,
- "p_value": p_value,
- "odds_ratio": odds_ratio,
- }
- )
- except:
- results.append(
- {
- "feature": feature,
- "comparison": f"{cluster1} vs {cluster2}",
- "chi2_statistic": np.nan,
- "p_value": np.nan,
- "odds_ratio": np.nan,
- }
- )
- return pd.DataFrame(results)
- def latex_crosstab(significance_df, gene_csv_path="raw/gene.csv"):
- gene_df = pd.read_csv(gene_csv_path)
- significance_df["Chromosome"] = significance_df["feature"].apply(
- lambda x: x.split(":")[0] if ":" in x else 0
- )
- significance_df["Position"] = significance_df["feature"].apply(
- lambda x: x.split(":")[1][:-1] if ":" in x else 0
- )
- significance_df["Genotype"] = significance_df["feature"].apply(
- lambda x: x.split(":")[1][-1] if ":" in x else 0
- )
- significance_df["Chromosome"] = significance_df["Chromosome"].astype(int)
- significance_df["Position"] = significance_df["Position"].astype(int)
- chrpos_rsID_dict = gene_df.set_index(["Chromosome", "Position"])["rsID"].to_dict()
- significance_df["rsID"] = significance_df.apply(
- lambda row: chrpos_rsID_dict.get((row["Chromosome"], row["Position"]), ""),
- axis=1,
- )
- significance_df = significance_df.merge(
- gene_df[["rsID", "Disease Type", "Ref", "Alt", "Gene"]], on="rsID", how="left"
- )
- significance_df = significance_df.sort_values(
- by=["comparison", "Disease Type", "Gene", "odds_ratio"],
- ascending=[True, True, True, False],
- )
- significance_df["Variant"] = significance_df["Ref"] + ">" + significance_df["Alt"]
- gene_filter = significance_df["feature"].str.startswith("E")
- significance_df.loc[gene_filter, "Variant"] = significance_df.loc[
- gene_filter, "feature"
- ]
- significance_df.loc[gene_filter, "Disease Type"] = "Dyslipidemia"
- significance_df.loc[gene_filter, "Gene"] = "APOE"
- significance_df.loc[gene_filter, "rsID"] = "rs429358, rs7412"
- significance_df["comparison"] = significance_df["comparison"].str.replace(
- "Reference vs ", "Ref. vs. "
- )
- disease_abbreviation = {"Dyslipidemia": "DL", "Type 2 Diabetes": "T2D"}
- significance_df["Disease Type"] = significance_df["Disease Type"].replace(
- disease_abbreviation
- )
- significance_df["rsID"] = significance_df["rsID"].str.replace("rs", "")
- gene_research_df = pd.read_csv("raw/gene_research.csv")
- gene_research_df["rsID"] = gene_research_df["rsID"].str.replace("rs", "")
- significance_df = significance_df.merge(
- gene_research_df[["rsID", "Risk Genotype"]], on="rsID", how="left"
- )
- e2f = significance_df["Variant"] == "E2"
- significance_df.loc[e2f, "Genotype"] = "TT"
- e3f = significance_df["Variant"] == "E3"
- significance_df.loc[e3f, "Genotype"] = "TC"
- e4f = significance_df["Variant"] == "E4"
- significance_df.loc[e4f, "Genotype"] = "CC"
- significance_df.loc[e2f, "Risk Genotype"] = "CC"
- significance_df.loc[e3f, "Risk Genotype"] = "CC"
- significance_df.loc[e4f, "Risk Genotype"] = "CC"
- return significance_df
- os.chdir("/workspace")
- num_folds = 5
- DATA_PATH = Path("data")
- DATA_PATH.mkdir(exist_ok=True, parents=True)
- df_total = pd.read_csv(DATA_PATH / "feature_fold0.csv")
- df_total[AGE_COLUMN_NAME] = pd.read_csv(DATA_PATH / "duration_fold0.csv")[
- AGE_COLUMN_NAME
- ]
- df_total[TARGET_COLUMN_NAME] = pd.read_csv(DATA_PATH / "event_fold0.csv")[
- TARGET_COLUMN_NAME
- ]
- df_total[FOLD_COLUMN_NAME] = 0
- for fold_idx in range(1, num_folds):
- df_fold = pd.read_csv(DATA_PATH / f"feature_fold{fold_idx}.csv")
- df_fold[AGE_COLUMN_NAME] = pd.read_csv(DATA_PATH / f"duration_fold{fold_idx}.csv")[
- AGE_COLUMN_NAME
- ]
- df_fold[TARGET_COLUMN_NAME] = pd.read_csv(DATA_PATH / f"event_fold{fold_idx}.csv")[
- TARGET_COLUMN_NAME
- ]
- df_fold[FOLD_COLUMN_NAME] = fold_idx
- df_total = pd.concat([df_total, df_fold], ignore_index=True)
- num_features = df_total.shape[1] - 3
- ckpts = []
- for fold_idx in range(num_folds):
- ckpt_dir = Path(f"ckpts/train/fold{fold_idx}")
- files = sorted([p for p in ckpt_dir.iterdir() if p.is_file()])
- if not files:
- raise FileNotFoundError(f"No files found in {ckpt_dir}")
- ckpt_path = files[0]
- ckpts.append(ckpt_path)
- # %%
- total_metrics = pd.DataFrame(
- columns=[
- "Fold",
- "FFN Train",
- "FFN Test",
- "LR Train",
- "LR Test",
- ]
- )
- for fold in range(num_folds):
- print(f"Fold {fold}")
- # Load the model state dict from the checkpoint
- deep_weibull_model = DeepWeibullModel(input_dim=num_features, aging_process=True)
- deep_weibull_model.load_state_dict(
- torch.load(ckpts[fold], map_location="cpu")["model"]
- )
- df_train = df_total[df_total[FOLD_COLUMN_NAME] != fold]
- df_test = df_total[df_total[FOLD_COLUMN_NAME] == fold]
- metric = get_fnn_result(
- df_train.drop(columns=[FOLD_COLUMN_NAME]),
- df_test.drop(columns=[FOLD_COLUMN_NAME]),
- deep_weibull_model,
- )
- c_index_fnn_train, c_index_fnn_test, _, _ = metric
- c_index_dict = {}
- c_index_dict["Fold"] = fold
- c_index_dict["FFN Train"] = c_index_fnn_train.item()
- c_index_dict["FFN Test"] = c_index_fnn_test.item()
- df_train_for_linear = df_train.drop(columns=[FOLD_COLUMN_NAME])
- df_train_for_linear[MMSE_COLUMN_NAME] = (
- df_train[MMSE_COLUMN_NAME].mean() / df_train[MMSE_COLUMN_NAME].std()
- )
- df_test_for_linear = df_test.drop(columns=[FOLD_COLUMN_NAME])
- df_test_for_linear[MMSE_COLUMN_NAME] = (
- df_test[MMSE_COLUMN_NAME].mean() / df_test[MMSE_COLUMN_NAME].std()
- )
- _, l_cindex_train, l_cindex_test = get_linear_aft_result(
- df_train_for_linear,
- df_test_for_linear,
- )
- c_index_dict["LR Train"] = l_cindex_train
- c_index_dict["LR Test"] = l_cindex_test
- total_metrics = pd.concat(
- [total_metrics, pd.DataFrame([c_index_dict])], ignore_index=True
- )
- mean_dict = {"Fold": "Mean"}
- std_dict = {"Fold": "Std"}
- for col in total_metrics.columns.drop("Fold"):
- mean_dict[col] = total_metrics[col].mean()
- std_dict[col] = total_metrics[col].std()
- total_metrics = pd.concat(
- [total_metrics, pd.DataFrame([mean_dict]), pd.DataFrame([std_dict])],
- ignore_index=True,
- )
- total_metrics.to_latex(
- "tables/c_index.tex",
- index=False,
- float_format="%.4f",
- caption="Concordance index(C-index) from 5-Fold Cross-Validation, LR: Linear Regression, FFN: Feedforward Neural Network, Std: Standard Deviation.",
- label="tab:c_index",
- position="htbp",
- )
- # %%
- gene_df = pd.read_csv("raw/gene.csv")
- chrpos_rsID = gene_df.set_index(["Chromosome", "Position"])["rsID"].to_dict()
- rename_map = {}
- used_names = set()
- for col in df_total.columns:
- if col == FOLD_COLUMN_NAME:
- continue
- new_name = col
- if ":" in col:
- try:
- chrom, rest = col.split(":", 1)
- chrom_i = int(chrom)
- pos = int(rest[:-1]) # drop allele letter at end
- rs = chrpos_rsID.get((chrom_i, pos))
- if isinstance(rs, str) and rs.strip():
- new_name = rs
- rename_map[col] = new_name + rest[-1]
- except Exception:
- new_name = col
- rename_map["E3"] = "APOE E3"
- rename_map["E4"] = "APOE E4"
- rename_map["E2"] = "APOE E2"
- # %%
- lr_summary, _, _ = get_linear_aft_result(
- df_total.drop(columns=[FOLD_COLUMN_NAME]), df_total.drop(columns=[FOLD_COLUMN_NAME])
- )
- lr_summary = (
- lr_summary.drop([("rho_", "Intercept")])
- .drop([("lambda_", "Intercept")])
- .reset_index()
- .drop(columns=["param"])
- )
- rename_map["E2"] = "APOE E2"
- rename_map["E3"] = "APOE E3"
- rename_map["E4"] = "APOE E4"
- rename_map["MMSE"] = "MMSE"
- rename_map["Female"] = "Female"
- lr_summary["covariate"] = lr_summary["covariate"].map(rename_map)
- lr_summary["abs_coef"] = lr_summary["coef"].abs()
- lr_summary = lr_summary[lr_summary["p"] < 0.05]
- lr_summary = lr_summary.sort_values(by=["p"], ascending=True)
- fig, ax = plt.subplots(figsize=(8, 4))
- lr_summary = lr_summary[:15]
- errors = [
- lr_summary["coef"] - lr_summary["coef lower 95%"],
- lr_summary["coef upper 95%"] - lr_summary["coef"],
- ]
- colors = ["red" if c < 0 else "royalblue" for c in lr_summary["coef"].values]
- ax.bar(
- x=lr_summary["covariate"],
- height=lr_summary["abs_coef"],
- color=colors,
- )
- ax.axhline(0, ls="--", color="gray", linewidth=2, zorder=0)
- ax.set_xticklabels(lr_summary["covariate"], rotation=90, ha="right")
- ax.set_ylabel("|Coefficient|")
- fig.savefig("fig_img/linear_aft_coef.pdf", dpi=300, bbox_inches="tight")
- # %%
- total_shap = pd.DataFrame()
- full_chpt = Path("ckpts/train/fold-1/checkpoint_-1.7854.pt")
- deep_weibull_model_full = DeepWeibullModel(
- input_dim=df_total.shape[1] - 3, aging_process=True
- )
- deep_weibull_model_full.load_state_dict(
- torch.load(full_chpt, map_location="cpu")["model"]
- )
- df_shap = df_total.rename(columns=rename_map)
- attributions_df = get_shap(
- model=deep_weibull_model_full,
- features=df_shap.drop(
- columns=[TARGET_COLUMN_NAME, AGE_COLUMN_NAME, FOLD_COLUMN_NAME]
- ),
- seed=42,
- )
- # %%
- top_features_shap = attributions_df.abs().mean().sort_values(ascending=False)
- top_n = 15
- fig, ax = plt.subplots(figsize=(8, 4))
- top15_features_shap = top_features_shap[:top_n]
- colors = [
- "red" if v < 0 else "royalblue"
- for v in attributions_df.rename(columns=rename_map)[top15_features_shap.index]
- .mean()
- .values
- ]
- ax.bar(
- x=top15_features_shap.index,
- height=top15_features_shap.values,
- # color=colors,
- )
- ax.set_xticklabels(top15_features_shap.index, rotation=90, ha="right")
- ax.set_ylabel("Mean(|SHAP value|)")
- fig.savefig("fig_img/shap_top15.pdf", dpi=300, bbox_inches="tight")
- # %%
- feature_counts = [10, 20, 30, 40]
- all_results = []
- df_total_for_subset = df_total.rename(columns=rename_map)
- for n in feature_counts:
- ckpt_dir = Path(f"ckpts/top{n}_train/fold-1")
- files = sorted([p for p in ckpt_dir.iterdir() if p.is_file()])
- if not files:
- raise FileNotFoundError(f"No checkpoint files found in {ckpt_dir}")
- chpt_path = files[0]
- top_features_shap_names = top_features_shap.index.tolist()[:n]
- model_n = DeepWeibullModel(input_dim=n, aging_process=True)
- model_n.load_state_dict(torch.load(chpt_path, map_location="cpu")["model"])
- df_subset = df_total_for_subset.drop(columns=[FOLD_COLUMN_NAME])[
- top_features_shap_names + [AGE_COLUMN_NAME, TARGET_COLUMN_NAME]
- ]
- c_index_train_n, c_index_test_n, brier_train_n, brier_test_n = get_fnn_result(
- df_subset, df_subset, model_n
- )
- all_results.append(
- {
- "Number of Features": n,
- "C-index": c_index_train_n.item(),
- "Brier": brier_train_n.item(),
- }
- )
- apoe_ckpt = Path("ckpts/top_apoe_train/fold-1/checkpoint_-1.8456.pt")
- model_apoe = DeepWeibullModel(input_dim=5, aging_process=True)
- model_apoe.load_state_dict(torch.load(apoe_ckpt, map_location="cpu")["model"])
- df_apoe = df_total.drop(columns=[FOLD_COLUMN_NAME])[
- ["E4", "MMSE", "E2", "E3", "Female", AGE_COLUMN_NAME, TARGET_COLUMN_NAME]
- ]
- (
- c_index_fnn_train_apoe,
- c_index_fnn_test_apoe,
- brier_fnn_train_apoe,
- brier_fnn_test_apoe,
- ) = get_fnn_result(df_apoe, df_apoe, model_apoe)
- all_results.append(
- {
- "Number of Features": 5,
- "C-index": c_index_fnn_train_apoe.item(),
- "Brier": brier_fnn_train_apoe.item(),
- }
- )
- full_chpt = Path("ckpts/train/fold-1/checkpoint_-1.7854.pt")
- deep_weibull_model_full = DeepWeibullModel(
- input_dim=df_total.shape[1] - 3, aging_process=True
- )
- deep_weibull_model_full.load_state_dict(
- torch.load(full_chpt, map_location="cpu")["model"]
- )
- (
- c_index_fnn_train_full,
- c_index_fnn_test_full,
- brier_fnn_train_full,
- brier_fnn_test_full,
- ) = get_fnn_result(
- df_total.drop(columns=[FOLD_COLUMN_NAME]),
- df_total.drop(columns=[FOLD_COLUMN_NAME]),
- deep_weibull_model_full,
- )
- all_results.append(
- {
- "Number of Features": df_total.shape[1] - 3,
- "C-index": c_index_fnn_train_full.item(),
- "Brier": brier_fnn_train_full.item(),
- }
- )
- all_cindex = pd.DataFrame(all_results).sort_values(by="Number of Features")
- fig, ax = plt.subplots(figsize=(6, 4))
- sns.lineplot(
- data=all_cindex,
- x="Number of Features",
- y="C-index",
- marker="o",
- ax=ax,
- )
- ax.set_xticks(all_cindex["Number of Features"])
- ax.set_ylim(0.65, 0.75)
- fig.savefig("fig_img/cindex_by_num_features.pdf", dpi=300, bbox_inches="tight")
- # %%
- fig, ax = plt.subplots(figsize=(6, 4))
- sns.lineplot(
- data=all_cindex.rename(columns={"Brier": "Integrated Brier Score"}),
- x="Number of Features",
- y="Integrated Brier Score",
- marker="o",
- ax=ax,
- )
- ax.set_xticks(all_cindex["Number of Features"])
- fig.savefig("fig_img/brier_by_num_features.pdf", dpi=300, bbox_inches="tight")
- # %%
- get_mult_shap_specific_interaction(
- model=deep_weibull_model_full,
- features=df_shap.drop(
- columns=[TARGET_COLUMN_NAME, AGE_COLUMN_NAME, FOLD_COLUMN_NAME]
- ),
- feature_xs=[
- "rs17145738T",
- "rs17145738G",
- "rs1558902A",
- "rs7007797G",
- "rs662799G",
- "rs9939609A",
- "rs708272A",
- ],
- feature_colors=["APOE E4", "APOE E3", "APOE E2"],
- seed=42,
- background_size=900,
- )
- # %%
visualize.py at commit b63c5e5, under MIT · at the source
Overview
- College of Pharmacy, Chungnam National University, Daejeon, 34134, Republic of Korea
- Department of Bio-AI convergence, Chungnam National University, Daejeon, 34134, Republic of Korea
- Institute of Drug Research and Development, Chungnam National University, Daejeon, 34134, Republic of Korea
- Department of Computer Science and Engineering, Chungnam National University, Daejeon, 34134, Republic of Korea
Abstract
Background: Alzheimer’s disease (AD) is a progressive neurodegenerative disease. Traditional models for estimating AD onset cannot capture nonlinear interactions (epistasis) among the numerous genetic variables that contribute to AD risk.
Methods: We developed a feedforward neural network (FFN)–Weibull survival model to predict AD onset using large-scale single-nucleotide polymorphism (SNP) data. We integrated an XAI technique, Shapley additive explanations (SHAP), to address the black-box nature of deep learning, interpret model predictions, and quantify the contribution of each genetic factor to AD.
Results: The FFN model achieved a mean concordance index of 0.647, demonstrating an approximately 3.6% improvement over the traditional linear baseline (0.625). The FFN-SHAP model validated established findings, identifying APOE E4 as a primary AD risk factor. APOE E2 strongly protected against AD. Metabolic-disorder-relat
Conclusions: By effectively bypassing the combinatorial explosion of interaction terms, the predictive power of an FFN combined with XAI provides a robust methodological tool for identifying the genetic basis of complex diseases, even in cohorts with limited sample sizes. Our model generated novel testable hypotheses regarding the intricate roles of gene–gene and gene–environment interactions in AD pathogenesis.
Reproduced under the paper's license (CC BY), from the paper cited above.
Repository
Its files are read in the Code ↔ Paper reader above, with 5 matches between paragraphs and lines of code.
swgoo/dementia_xai_sa
b63c5e5391cf9927db947aa4af0f4a5991e9e29f, 24 April 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
3 files
- model.py — Python, 352 lines
- visualize.py — Python, 888 lines, 5 matches
- LICENSE — License, 21 lines
Code availability
The source codes for the proposed FFN-Weibull model and SHAP analysis are provided at git repository (https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 2 scripts, each with its path and the digest of its content;
- 5 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
No dataset and no data link were found in the paper.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 6 MeSH terms, 20 funders, 32 references.
Cite
This paper
Goo, S., Lee, S., Chae, J.-W., Jung, S., & Yun, H.-y. (2026). Interpretable deep survival analysis of Alzheimer's disease via metabolic genetic variants. Bioinformatics (Oxford, England), 42(6), btag213. https://
BibTeX
@article{goo2026interpre
author = {Goo, Sungwoo and Lee, Soyoung and Chae, Jung-Woo and Jung, Sangkeun and Yun, Hwi-yeol},
title = {{Interpretable deep survival analysis of Alzheimer's disease via metabolic genetic variants}},
journal = {Bioinformatics (Oxford, England)},
year = {2026},
month = jun,
volume = {42},
number = {6},
pages = {btag213},
publisher = {Oxford University Press},
issn = {1367-4803},
doi = {10.1093/
url = {https://
pmid = {42063212},
pmcid = {PMC13224968}
}
RIS
TY - JOUR
AU - Goo, Sungwoo
AU - Lee, Soyoung
AU - Chae, Jung-Woo
AU - Jung, Sangkeun
AU - Yun, Hwi-yeol
TI - Interpretable deep survival analysis of Alzheimer's disease via metabolic genetic variants
T2 - Bioinformatics (Oxford, England)
J2 - Bioinformatics
PY - 2026
DA - 2026/
VL - 42
IS - 6
SP - btag213
SN - 1367-4803
PB - Oxford University Press
DO - 10.1093/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1093/
"type": "article-journal",
"title": "Interpretable deep survival analysis of Alzheimer's disease via metabolic genetic variants",
"container-title": "Bioinformatics (Oxford, England)",
"author": [
{
"family": "Goo",
"given": "Sungwoo"
},
{
"family": "Lee",
"given": "Soyoung"
},
{
"family": "Chae",
"given": "Jung-Woo"
},
{
"family": "Jung",
"given": "Sangkeun"
},
{
"family": "Yun",
"given": "Hwi-yeol"
}
],
"container-title-short":
"volume": "42",
"issue": "6",
"page": "btag213",
"DOI": "10.1093/
"PMID": "42063212",
"PMCID": "PMC13224968",
"ISSN": "1367-4803",
"publisher": "Oxford University Press",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
1
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1038/s41598-026-48613-0 [code]
- An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging.Journal: Scientific reportsIn common: SHAP, PyTorch, seaborn, 5 other tools, genetics / omics, 1 reference
- [2] doi:10.1038/s41467-026-76837-1 [code]
- Drug screen and machine learning predict neuroprotective agents in a preclinical human model of childhood dementia.Journal: Nature communicationsIn common: SHAP, PyTorch, seaborn, 5 other tools, Alzheimer's / dementia
- [3] doi:10.1038/s43856-026-01606-6 [code]
- Validation of remote multimodal AI screening for Parkinson disease across diverse settings.Journal: Communications medicineIn common: SHAP, PyTorch, seaborn, 5 other tools, clinical / translational
- [4] doi:10.1038/s41467-026-71555-0 [code]
- A deep representation learning model to predict response to vagus nerve stimulation.Journal: Nature communicationsIn common: SHAP, PyTorch, seaborn, 5 other tools, clinical / translational
- [5] doi:10.1186/s13059-026-04125-8 [code]
- MLMarker: a machine learning framework for tissue inference and biomarker discovery.Journal: Genome biologyIn common: SHAP, PyTorch, seaborn, 5 other tools, genetics / omics
- [6] doi:10.1016/j.isci.2026.116825 [code]
- Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.Journal: iScienceIn common: SHAP, PyTorch, seaborn, 5 other tools, genetics / omics
- [7] doi:10.1038/s42003-026-10957-8 [code]
- Brain defence by the extracellular matrix protein Cochlin.Journal: Communications biologyIn common: SHAP, PyTorch, seaborn, 5 other tools
- [8] doi:10.1371/journal.pone.0345854 [code]
- Shedding light on neural learning to rank models for anticancer drug prioritization.Journal: PloS oneIn common: SHAP, PyTorch, seaborn, 5 other tools
- [9] doi:10.1371/journal.pcbi.1014615 [code]
- Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains.Journal: PLoS computational biologyIn common: SHAP, PyTorch, seaborn, 5 other tools
- [10] doi:10.1093/nargab/lqag050 [code]
- TSProm: deep learning framework to predict tissue-specific regulatory logic.Journal: NAR genomics and bioinformaticsIn common: SHAP, PyTorch, seaborn, 5 other tools
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 2 scripts, and 5 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:105a3f55634f0bea…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
