diff --git a/notebooks/kaplan_meier.ipynb b/notebooks/kaplan_meier.ipynb index b882e7b..a1a5df9 100644 --- a/notebooks/kaplan_meier.ipynb +++ b/notebooks/kaplan_meier.ipynb @@ -10,7 +10,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": 7, "id": "075dad06", "metadata": {}, "outputs": [ @@ -30,7 +30,7 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": 8, "id": "6649f9d0", "metadata": {}, "outputs": [], @@ -53,7 +53,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": 9, "id": "09c2011b", "metadata": {}, "outputs": [ @@ -124,14 +124,14 @@ }, { "cell_type": "code", - "execution_count": 9, + "execution_count": 10, "id": "f7c08980", "metadata": {}, "outputs": [], "source": [ "import cenreg.model.nonparametric\n", "\n", - "dist = cenreg.model.nonparametric.kaplan_meier_estimator(df[\"time\"].values, df[\"event\"].values)" + "dist = cenreg.model.nonparametric.kaplan_meier_estimator(df[\"time\"].values.astype(float), df[\"event\"].values)" ] }, { @@ -144,7 +144,7 @@ }, { "cell_type": "code", - "execution_count": 10, + "execution_count": 11, "id": "7aebf393", "metadata": {}, "outputs": [ diff --git a/notebooks/li_watkins_yu.ipynb b/notebooks/li_watkins_yu.ipynb index 3653856..ac15cd4 100644 --- a/notebooks/li_watkins_yu.ipynb +++ b/notebooks/li_watkins_yu.ipynb @@ -10,19 +10,10 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": 1, "id": "458fd7ec", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "The autoreload extension is already loaded. To reload it, use:\n", - " %reload_ext autoreload\n" - ] - } - ], + "outputs": [], "source": [ "%load_ext autoreload\n", "%autoreload 2" @@ -38,7 +29,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": 2, "id": "1615397c", "metadata": {}, "outputs": [], @@ -59,7 +50,7 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": 3, "id": "bf0a7543", "metadata": {}, "outputs": [], @@ -79,7 +70,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": 4, "id": "8c73dfeb", "metadata": {}, "outputs": [ diff --git a/notebooks/nn_ic_log.ipynb b/notebooks/nn_ic_log.ipynb index 6898a10..18969a5 100644 --- a/notebooks/nn_ic_log.ipynb +++ b/notebooks/nn_ic_log.ipynb @@ -10,19 +10,10 @@ }, { "cell_type": "code", - "execution_count": 103, + "execution_count": 1, "id": "34da57fc", "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "The autoreload extension is already loaded. To reload it, use:\n", - " %reload_ext autoreload\n" - ] - } - ], + "outputs": [], "source": [ "%load_ext autoreload\n", "%autoreload 2" @@ -30,7 +21,7 @@ }, { "cell_type": "code", - "execution_count": 104, + "execution_count": 2, "id": "0371fda9", "metadata": {}, "outputs": [], @@ -57,7 +48,7 @@ }, { "cell_type": "code", - "execution_count": 105, + "execution_count": 3, "id": "cabba736", "metadata": {}, "outputs": [ @@ -101,7 +92,7 @@ }, { "cell_type": "code", - "execution_count": 106, + "execution_count": 4, "id": "1385c74b", "metadata": {}, "outputs": [ @@ -155,7 +146,7 @@ }, { "cell_type": "code", - "execution_count": 107, + "execution_count": 5, "id": "4c810bf4", "metadata": {}, "outputs": [ @@ -215,7 +206,7 @@ }, { "cell_type": "code", - "execution_count": 108, + "execution_count": 6, "id": "65de15bf", "metadata": {}, "outputs": [ @@ -247,7 +238,7 @@ }, { "cell_type": "code", - "execution_count": 109, + "execution_count": 7, "id": "d57af800", "metadata": {}, "outputs": [], @@ -281,7 +272,7 @@ }, { "cell_type": "code", - "execution_count": 110, + "execution_count": 8, "id": "596ac944", "metadata": {}, "outputs": [ @@ -330,7 +321,7 @@ }, { "cell_type": "code", - "execution_count": 111, + "execution_count": 9, "id": "2c9d35a4", "metadata": {}, "outputs": [], @@ -361,7 +352,7 @@ }, { "cell_type": "code", - "execution_count": 114, + "execution_count": 10, "id": "a7f84c10", "metadata": {}, "outputs": [ diff --git a/notebooks/sc_net.ipynb b/notebooks/sc_net.ipynb index a6b38f2..16a2160 100644 --- a/notebooks/sc_net.ipynb +++ b/notebooks/sc_net.ipynb @@ -350,7 +350,7 @@ }, { "cell_type": "code", - "execution_count": 12, + "execution_count": 10, "id": "b050886e", "metadata": {}, "outputs": [ @@ -361,20 +361,10 @@ "CJD-Brier 0.9616783072036734\n", "CJD-Logarithmic 3.6927649835189023\n", "CJD-KS 0.2863711858453328\n", - "NLL-SC 4.885996452186602\n", - "Cen-log 0.4838524\n", - "D-calibration 0.7879467415784499\n" - ] - }, - { - "ename": "AttributeError", - "evalue": "module 'cenreg.metric.cdf' has no attribute 'km_calibration'", - "output_type": "error", - "traceback": [ - "\u001b[31m---------------------------------------------------------------------------\u001b[39m", - "\u001b[31mAttributeError\u001b[39m Traceback (most recent call last)", - "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[12]\u001b[39m\u001b[32m, line 38\u001b[39m\n\u001b[32m 35\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33m\"\u001b[39m\u001b[33mD-calibration\u001b[39m\u001b[33m\"\u001b[39m, dcal)\n\u001b[32m 37\u001b[39m \u001b[38;5;66;03m# Compute KM-calibration\u001b[39;00m\n\u001b[32m---> \u001b[39m\u001b[32m38\u001b[39m kmcal = \u001b[43mcenreg\u001b[49m\u001b[43m.\u001b[49m\u001b[43mmetric\u001b[49m\u001b[43m.\u001b[49m\u001b[43mcdf\u001b[49m\u001b[43m.\u001b[49m\u001b[43mkm_calibration\u001b[49m(list_t_dist[\u001b[32m1\u001b[39m], observed_times, events, bins_np)\n\u001b[32m 39\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33m\"\u001b[39m\u001b[33mKM-calibration\u001b[39m\u001b[33m\"\u001b[39m, kmcal)\n", - "\u001b[31mAttributeError\u001b[39m: module 'cenreg.metric.cdf' has no attribute 'km_calibration'" + "NLL-SC 4.8859964377774485\n", + "Cen-log 0.37202982911557014\n", + "D-calibration 0.0005720768376036918\n", + "KM-calibration 3.5778655806086626\n" ] } ], @@ -422,7 +412,7 @@ }, { "cell_type": "code", - "execution_count": null, + "execution_count": 11, "id": "543f50f2", "metadata": {}, "outputs": [ diff --git a/notebooks/ts_brier.ipynb b/notebooks/ts_brier.ipynb index 02443a4..84f411f 100644 --- a/notebooks/ts_brier.ipynb +++ b/notebooks/ts_brier.ipynb @@ -10,10 +10,19 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": 12, "id": "5d7b84bc", "metadata": {}, - "outputs": [], + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "The autoreload extension is already loaded. To reload it, use:\n", + " %reload_ext autoreload\n" + ] + } + ], "source": [ "%load_ext autoreload\n", "%autoreload 2" @@ -21,7 +30,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": 13, "id": "ae19496a", "metadata": {}, "outputs": [], @@ -48,7 +57,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": 14, "id": "d9d1e139", "metadata": {}, "outputs": [ @@ -108,7 +117,7 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": 15, "id": "b4f78464", "metadata": {}, "outputs": [ @@ -185,7 +194,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": 16, "id": "c8ba9e12", "metadata": {}, "outputs": [ @@ -214,7 +223,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": 17, "id": "c4d5ce80", "metadata": {}, "outputs": [], @@ -230,16 +239,16 @@ "\n", "# preparation for PyTorch training\n", "sdm = SurvDataModule(128)\n", - "train_dataloader = sdm.train_dataloader(x_train, df_train[\"time\"].values, df_train[\"event\"].values)\n", + "train_dataloader = sdm.train_dataloader(x_train, df_train[\"time\"].values.astype(float), df_train[\"event\"].values)\n", "num_bins = 32\n", - "bins = torch.tensor(cenreg.utils.create_bins(df[\"time\"].max(), 0.0, num_bins))\n", + "bins = torch.tensor(cenreg.utils.create_bins(df[\"time\"].max().astype(float), 0.0, num_bins))\n", "loss_fn = Brier(bins, 2)\n", "model = MLP(x_train.shape[1], num_bins*2, 64)" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": 18, "id": "2e5f8205", "metadata": {}, "outputs": [ @@ -285,7 +294,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": 19, "id": "3429b83c", "metadata": {}, "outputs": [ @@ -321,7 +330,7 @@ }, { "cell_type": "code", - "execution_count": 9, + "execution_count": 20, "id": "62b4f42e", "metadata": {}, "outputs": [ @@ -390,7 +399,7 @@ }, { "cell_type": "code", - "execution_count": 10, + "execution_count": 21, "id": "a3cbdc55", "metadata": {}, "outputs": [ @@ -401,7 +410,7 @@ "CJD-Brier 0.6394180830320167\n", "CJD-Logarithmic 2.379232\n", "CJD-KS 0.1807599546171742\n", - "NLL-SC 6.258204767185328\n" + "NLL-SC 6.533535194080549\n" ] } ], @@ -411,7 +420,7 @@ "import cenreg.model.copula_np\n", "\n", "cjd_pred = cjd_pred.reshape(-1, num_bins * 2)\n", - "observed_times = df_test['time'].values\n", + "observed_times = df_test['time'].values.astype(float)\n", "events = df_test[\"event\"].astype(bool).values\n", "bins_np = bins.detach().cpu().numpy()\n", "\n", @@ -436,7 +445,7 @@ }, { "cell_type": "code", - "execution_count": 11, + "execution_count": 22, "id": "5a6a98f6", "metadata": {}, "outputs": [ diff --git a/notebooks/ts_lgb.ipynb b/notebooks/ts_lgb.ipynb index ceab3b5..6936491 100644 --- a/notebooks/ts_lgb.ipynb +++ b/notebooks/ts_lgb.ipynb @@ -11,10 +11,19 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": 14, "id": "a879879e", "metadata": {}, - "outputs": [], + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "The autoreload extension is already loaded. To reload it, use:\n", + " %reload_ext autoreload\n" + ] + } + ], "source": [ "%load_ext autoreload\n", "%autoreload 2" @@ -22,7 +31,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": 15, "id": "0e3cb25d", "metadata": {}, "outputs": [], @@ -45,7 +54,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": 16, "id": "abbbdc13", "metadata": {}, "outputs": [ @@ -116,7 +125,7 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": 17, "id": "b1a07ae2", "metadata": {}, "outputs": [ @@ -159,7 +168,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": 18, "id": "fa2dd5a8", "metadata": {}, "outputs": [ @@ -213,7 +222,7 @@ }, { "cell_type": "code", - "execution_count": 6, + "execution_count": 19, "id": "ad343a33", "metadata": {}, "outputs": [ @@ -244,7 +253,7 @@ }, { "cell_type": "code", - "execution_count": 7, + "execution_count": 20, "id": "1006fee1", "metadata": {}, "outputs": [ @@ -300,7 +309,7 @@ }, { "cell_type": "code", - "execution_count": 8, + "execution_count": 21, "id": "3e7ec2fb", "metadata": {}, "outputs": [ @@ -373,7 +382,7 @@ }, { "cell_type": "code", - "execution_count": 11, + "execution_count": 22, "id": "8a8612a2", "metadata": {}, "outputs": [ @@ -381,23 +390,14 @@ "name": "stdout", "output_type": "stream", "text": [ - "CJD-Brier 0.7090305785585743\n", - "CJD-Logarithmic 2.0910114724468554\n", - "CJD-KS 0.22946618670525376\n", - "NLL-SC 7.1811389332790885\n", - "Cen-log 1.0275933871938163\n", - "D-calibration 0.4382083144368859\n" - ] - }, - { - "ename": "AttributeError", - "evalue": "module 'cenreg.metric.cdf' has no attribute 'km_calibration'", - "output_type": "error", - "traceback": [ - "\u001b[31m---------------------------------------------------------------------------\u001b[39m", - "\u001b[31mAttributeError\u001b[39m Traceback (most recent call last)", - "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[11]\u001b[39m\u001b[32m, line 42\u001b[39m\n\u001b[32m 39\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33m\"\u001b[39m\u001b[33mD-calibration\u001b[39m\u001b[33m\"\u001b[39m, dcal)\n\u001b[32m 41\u001b[39m \u001b[38;5;66;03m# Compute KM-calibration\u001b[39;00m\n\u001b[32m---> \u001b[39m\u001b[32m42\u001b[39m kmcal = \u001b[43mcenreg\u001b[49m\u001b[43m.\u001b[49m\u001b[43mmetric\u001b[49m\u001b[43m.\u001b[49m\u001b[43mcdf\u001b[49m\u001b[43m.\u001b[49m\u001b[43mkm_calibration\u001b[49m(list_t_dist[\u001b[32m1\u001b[39m], observed_times, events, bins)\n\u001b[32m 43\u001b[39m \u001b[38;5;28mprint\u001b[39m(\u001b[33m\"\u001b[39m\u001b[33mKM-calibration\u001b[39m\u001b[33m\"\u001b[39m, kmcal)\n", - "\u001b[31mAttributeError\u001b[39m: module 'cenreg.metric.cdf' has no attribute 'km_calibration'" + "CJD-Brier: 0.7090305785585743\n", + "CJD-Logarithmic: 2.0910114724468554\n", + "CJD-KS: 0.22946618670525376\n", + "NLL-SC: 7.1811389332790885\n", + "Cen-log: 0.9808302589716417\n", + "D-calibration (with 10 bins): 0.00205484143204613\n", + "KM-calibration (with 32 bins): 0.022774582695642684\n", + "IC-Cal: 0.0005136724911047157\n" ] } ], @@ -417,39 +417,46 @@ "\n", "# Compute Brier score on CJD representation\n", "cjd_brier = cenreg.metric.cjd.brier(observed_times, events, 2, cjd_pred, bins)\n", - "print(\"CJD-Brier\", cjd_brier.mean())\n", + "print(\"CJD-Brier:\", cjd_brier.mean())\n", "\n", "# Compute Logarithmic score on CDF representation\n", "cjd_logarithmic = cenreg.metric.cjd.negative_loglikelihood(observed_times, events, cjd_pred, bins)\n", - "print(\"CJD-Logarithmic\", cjd_logarithmic.mean())\n", + "print(\"CJD-Logarithmic:\", cjd_logarithmic.mean())\n", "\n", "# Compute KS calibration error on CDF representation\n", "cjd_ks = cenreg.metric.cjd.kolmogorov_smirnov_calibration_error(observed_times, events, cjd_pred, bins)\n", - "print(\"CJD-KS\", cjd_ks)\n", + "print(\"CJD-KS:\", cjd_ks)\n", "\n", "# Compute NLL-SC metric\n", "copula_np = cenreg.model.copula_np.IndependenceCopula()\n", "survival_copula_np = cenreg.model.copula_np.SurvivalCopula(copula_np)\n", "nll_sc = cenreg.metric.cdf.nll_sc(list_t_dist, observed_times, events, survival_copula_np)\n", - "print(\"NLL-SC\", nll_sc.mean())\n", + "print(\"NLL-SC:\", nll_sc.mean())\n", "\n", "# Compute cen-log metric\n", "nll = cenreg.metric.cdf.negative_log_likelihood_survival(list_t_dist[1], observed_times, events)\n", - "print(\"Cen-log\", nll.mean())\n", + "print(\"Cen-log:\", nll.mean())\n", "\n", "# Compute D-calibration\n", "quantiles_cal = np.linspace(0.0, 1.0, 11)\n", "dcal = cenreg.metric.quantile.d_calibration(list_t_dist[1], observed_times, events, quantiles_cal)\n", - "print(\"D-calibration\", dcal)\n", + "print(f\"D-calibration (with {len(quantiles_cal)-1} bins):\", dcal)\n", "\n", "# Compute KM-calibration\n", "kmcal = cenreg.metric.cdf.km_calibration(list_t_dist[1], observed_times, events, bins)\n", - "print(\"KM-calibration\", kmcal)\n" + "print(f\"KM-calibration (with {len(bins)-1} bins):\", kmcal)\n", + "\n", + "# Compute IC-Log calibration error on CDF representation\n", + "lb_test = observed_times\n", + "ub_test = np.full(observed_times.shape, np.inf)\n", + "ub_test[events] = observed_times[events] + 0.0001\n", + "ic_cal = cenreg.metric.quantile.ic_calibration(list_t_dist[1], lb_test, ub_test)\n", + "print(\"IC-Cal:\", ic_cal)" ] }, { "cell_type": "code", - "execution_count": null, + "execution_count": 23, "id": "8fca2660", "metadata": {}, "outputs": [ diff --git a/notebooks/zheng_klein.ipynb b/notebooks/zheng_klein.ipynb index 0bef4d1..7334627 100644 --- a/notebooks/zheng_klein.ipynb +++ b/notebooks/zheng_klein.ipynb @@ -10,10 +10,19 @@ }, { "cell_type": "code", - "execution_count": 1, + "execution_count": 6, "id": "af4e7383", "metadata": {}, - "outputs": [], + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "The autoreload extension is already loaded. To reload it, use:\n", + " %reload_ext autoreload\n" + ] + } + ], "source": [ "%load_ext autoreload\n", "%autoreload 2" @@ -21,7 +30,7 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": 7, "id": "dd4921a8", "metadata": {}, "outputs": [], @@ -44,7 +53,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": 8, "id": "5cb6ab77", "metadata": {}, "outputs": [ @@ -115,7 +124,7 @@ }, { "cell_type": "code", - "execution_count": 4, + "execution_count": 9, "id": "f6042a93", "metadata": {}, "outputs": [], @@ -126,11 +135,12 @@ "copula_ub = FrankCopula(5.0)\n", "copula_lb = FrankCopula(-5.0)\n", "copula_ind = IndependenceCopula()\n", - "dist_ind = cenreg.model.nonparametric.zheng_klein_estimator(df[\"time\"].values, df[\"event\"].values, copula_ind)\n", + "observed_times = df[\"time\"].values.astype(float)\n", + "dist_ind = cenreg.model.nonparametric.zheng_klein_estimator(observed_times, df[\"event\"].values, copula_ind)\n", "\n", "# Obtain upper and lower bound distributions under copula uncertainty\n", - "dist_ub = cenreg.model.nonparametric.zheng_klein_estimator(df[\"time\"].values, df[\"event\"].values, copula_ub)\n", - "dist_lb = cenreg.model.nonparametric.zheng_klein_estimator(df[\"time\"].values, df[\"event\"].values, copula_lb)" + "dist_ub = cenreg.model.nonparametric.zheng_klein_estimator(observed_times, df[\"event\"].values, copula_ub)\n", + "dist_lb = cenreg.model.nonparametric.zheng_klein_estimator(observed_times, df[\"event\"].values, copula_lb)" ] }, { @@ -143,7 +153,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": 10, "id": "ab1868b7", "metadata": {}, "outputs": [ diff --git a/src/cenreg/distribution/interpolate.py b/src/cenreg/distribution/interpolate.py index e4bbc3a..75e853c 100644 --- a/src/cenreg/distribution/interpolate.py +++ b/src/cenreg/distribution/interpolate.py @@ -25,7 +25,8 @@ def linear_interpolation(kx: np.ndarray, ky: np.ndarray, x: np.ndarray) -> np.nd if kx.ndim == 1: if ky.ndim == 2: - assert kx.shape[0] == ky.shape[1] + if kx.shape[0] != ky.shape[1]: + raise ValueError("kx and ky must have compatible shapes") # compute idx and ratio if kx.ndim == 1: diff --git a/src/cenreg/metric/cdf.py b/src/cenreg/metric/cdf.py index 06a9086..510af2f 100644 --- a/src/cenreg/metric/cdf.py +++ b/src/cenreg/metric/cdf.py @@ -1,3 +1,5 @@ +import warnings + import numpy as np from cenreg.model.nonparametric import kaplan_meier_estimator @@ -271,8 +273,10 @@ def brier( loss : Array of shape [batch_size] """ - assert len(y.shape) == 1 - assert len(y_bins.shape) == 1 + if len(y.shape) != 1: + raise ValueError("y must be of shape [batch_size]") + if len(y_bins.shape) != 1: + raise ValueError("y_bins must be of shape [num_bins+1]") F_pred = dist.cdf(y_bins) if len(F_pred.shape) == 1: @@ -306,8 +310,10 @@ def ranked_probability_score( loss : Array of shape [batch_size] """ - assert len(y.shape) == 1 - assert len(y_bins.shape) == 1 + if len(y.shape) != 1: + raise ValueError("y must be of shape [batch_size]") + if len(y_bins.shape) != 1: + raise ValueError("y_bins must be of shape [num_col]") F_pred = dist.cdf(y_bins[1:-1]) y = y.reshape(-1, 1) @@ -329,9 +335,15 @@ def nll_sc( Compute Negative Log-Likelihood based on Survival Copula (NLL-SC). """ - assert len(observed_times.shape) == 1 - assert len(events.shape) == 1 - assert observed_times.shape[0] == events.shape[0] + if len(observed_times.shape) != 1: + raise ValueError("observed_times must be of shape [batch_size]") + if len(events.shape) != 1: + raise ValueError("events must be of shape [batch_size]") + if observed_times.shape[0] != events.shape[0]: + raise ValueError("observed_times and events must have the same length") + if not np.issubdtype(observed_times.dtype, np.floating): + warnings.warn("observed_times is not float, converting to float.", stacklevel=2) + observed_times = observed_times.astype(float) Sl_list = [] Sr_list = [] diff --git a/src/cenreg/metric/cjd.py b/src/cenreg/metric/cjd.py index 7ef7382..5681f61 100644 --- a/src/cenreg/metric/cjd.py +++ b/src/cenreg/metric/cjd.py @@ -162,11 +162,16 @@ def kolmogorov_smirnov_calibration_error( Sum of Kolmogorov-Sminov calibration error. """ - assert len(observed_times.shape) == 1 - assert len(events.shape) == 1 - assert len(f_pred.shape) == 2 - assert len(boundaries.shape) == 1 - assert f_pred.shape[0] == observed_times.shape[0] + if len(observed_times.shape) != 1: + raise ValueError("observed_times must be of shape [batch_size]") + if len(events.shape) != 1: + raise ValueError("events must be of shape [batch_size]") + if len(f_pred.shape) != 2: + raise ValueError("f_pred must be of shape [batch_size, num_bin*num_risks]") + if len(boundaries.shape) != 1: + raise ValueError("boundaries must be of shape [num_bin+1]") + if f_pred.shape[0] != observed_times.shape[0]: + raise ValueError("f_pred and observed_times must have the same length") events = events.astype(int).reshape(-1, 1) idx = np.searchsorted(boundaries, observed_times.reshape(-1, 1), side="right") diff --git a/src/cenreg/metric/quantile.py b/src/cenreg/metric/quantile.py index cd6e816..24b5ea8 100644 --- a/src/cenreg/metric/quantile.py +++ b/src/cenreg/metric/quantile.py @@ -95,9 +95,12 @@ def ic_calibration( Value of IC-Cal. """ - assert len(lb.shape) == 1 - assert len(ub.shape) == 1 - assert lb.shape[0] == ub.shape[0] + if len(lb.shape) != 1: + raise ValueError("lb must be of shape [batch_size]") + if len(ub.shape) != 1: + raise ValueError("ub must be of shape [batch_size]") + if lb.shape[0] != ub.shape[0]: + raise ValueError("lb and ub must have the same length") if p != 2.0: raise NotImplementedError("Only p=2 is implemented.") diff --git a/src/cenreg/model/cjd2F_np.py b/src/cenreg/model/cjd2F_np.py index ae5f405..061eac6 100644 --- a/src/cenreg/model/cjd2F_np.py +++ b/src/cenreg/model/cjd2F_np.py @@ -16,7 +16,8 @@ def _integral(jd_pred: np.ndarray) -> np.ndarray: F_pred: estimated CDF. np.ndarray of shape [batch_size, num_risks, num_bin_predictions+1] """ - assert len(jd_pred.shape) == 3 + if len(jd_pred.shape) != 3: + raise ValueError(f"Expected jd_pred to have 3 dimensions, but got {len(jd_pred.shape)} dimensions.") # w = boundaries[1:] - boundaries[:-1] Q = np.cumsum(jd_pred, axis=2) diff --git a/src/cenreg/model/copula_np.py b/src/cenreg/model/copula_np.py index 2089aee..2541aab 100644 --- a/src/cenreg/model/copula_np.py +++ b/src/cenreg/model/copula_np.py @@ -25,7 +25,10 @@ def cdf(self, u: np.ndarray) -> np.ndarray: probability : ndarray (float) ndarray of shape [batch_size]. """ - assert u.ndim == 2 and u.shape[1] == 2, "Input must be a 2D array with shape [batch_size, 2]" + if u.ndim != 2: + raise ValueError("u must be 2-dimensional array.") + if u.shape[1] != 2: + raise ValueError("u must have shape [batch_size, 2].") return np.prod(u, axis=1) diff --git a/src/cenreg/model/nonparametric.py b/src/cenreg/model/nonparametric.py index bb5e265..da492d6 100644 --- a/src/cenreg/model/nonparametric.py +++ b/src/cenreg/model/nonparametric.py @@ -1,3 +1,5 @@ +import warnings + import numpy as np from cenreg.distribution.cdf import CumulativeDist @@ -32,6 +34,17 @@ def _set_ymin_ymax( return temp_min, temp_max +def _validate_weights(weights: np.ndarray, y: np.ndarray): + if weights is None: + return np.ones_like(y) + if len(weights.shape) != 1: + raise ValueError("weight must be one-dimensional array.") + if weights.shape[0] != y.shape[0]: + raise ValueError("weight and y must have the same length.") + if np.any(weights < 0.0): + raise ValueError("weight must be non-negative.") + + def _validate_cdf_inputs( y: np.ndarray, weights: np.ndarray | None = None, @@ -42,17 +55,14 @@ def _validate_cdf_inputs( raise ValueError("y must be one-dimensional array.") if y.size == 0: raise ValueError("y must not be empty.") - if weights is not None: - if len(weights.shape) != 1: - raise ValueError("weight must be one-dimensional array.") - if weights.shape[0] != y.shape[0]: - raise ValueError("weight and y must have the same length.") - if np.any(weights < 0.0): - raise ValueError("weight must be non-negative.") if y_min is not None: - assert y_min <= np.min(y), "y_min must be less than or equal to min(y)." + if y_min > np.min(y): + raise ValueError("y_min must be less than or equal to min(y).") if y_max is not None: - assert y_max >= np.max(y), "y_max must be greater than or equal to max(y)." + if y_max < np.max(y): + raise ValueError("y_max must be greater than or equal to max(y).") + if weights is not None: + _validate_weights(weights, y) def _adjust_bins(bins: np.ndarray, y_min: float, y_max: float) -> np.ndarray: @@ -122,6 +132,36 @@ def empirical_cdf_estimator( return CumulativeDist(b=bins, cum_p=cum_p, interpolate="right") +def _validate_kaplan_meier_inputs_weights(weights: np.ndarray, observed_times: np.ndarray): + if len(weights.shape) != 1: + raise ValueError("weights must be one-dimensional array.") + if observed_times.shape[0] != weights.shape[0]: + raise ValueError("observed_times and weights must have the same length.") + + +def _validate_kaplan_meier_inputs( + observed_times: np.ndarray, + uncensored: np.ndarray, + weights: np.ndarray | None = None, + y_min: float | None = None, + y_max: float | None = None, +): + if len(observed_times.shape) != 1: + raise ValueError("observed_times must be one-dimensional array.") + if len(uncensored.shape) != 1: + raise ValueError("uncensored must be one-dimensional array.") + if observed_times.shape[0] != uncensored.shape[0]: + raise ValueError("observed_times and uncensored must have the same length.") + if y_min is not None: + if y_min > np.min(observed_times): + raise ValueError("y_min must be less than or equal to min(observed_times).") + if y_max is not None: + if y_max < np.max(observed_times): + raise ValueError("y_max must be greater than or equal to max(observed_times).") + if weights is not None: + _validate_kaplan_meier_inputs_weights(weights, observed_times) + + def kaplan_meier_estimator( observed_times: np.ndarray, uncensored: np.ndarray, @@ -151,21 +191,16 @@ def kaplan_meier_estimator( Cumulative distribution function. """ - assert len(observed_times.shape) == 1 - assert len(uncensored.shape) == 1 - assert observed_times.shape[0] == uncensored.shape[0] + if not np.issubdtype(observed_times.dtype, np.floating): + warnings.warn("observed_times is not float, converting to float.", stacklevel=2) + observed_times = observed_times.astype(float) + _validate_kaplan_meier_inputs(observed_times, uncensored, weights, y_min, y_max) + uncensored = uncensored.astype(int) if np.sum(uncensored) == 0: raise ValueError("At least one data point must be uncensored.") if weights is None: weights = np.ones_like(observed_times) - else: - assert len(weights.shape) == 1 - assert observed_times.shape[0] == weights.shape[0] - if y_min is not None: - assert y_min <= np.min(observed_times), "y_min must be less than or equal to min(observed_times)." - if y_max is not None: - assert y_max >= np.max(observed_times), "y_max must be greater than or equal to max(observed_times)." # sort based on uncensored and observed_times temp = np.concatenate( @@ -312,17 +347,29 @@ def zheng_klein_estimator( Cumulative distribution function. """ - assert len(observed_times.shape) == 1 - assert len(uncensored.shape) == 1 - assert observed_times.shape[0] == uncensored.shape[0] + if len(observed_times.shape) != 1: + raise ValueError("observed_times must be one-dimensional array.") + if len(uncensored.shape) != 1: + raise ValueError("uncensored must be one-dimensional array.") + if observed_times.shape[0] != uncensored.shape[0]: + raise ValueError("observed_times and uncensored must have the same length.") + if not np.issubdtype(observed_times.dtype, np.floating): + warnings.warn("observed_times is not float, converting to float.", stacklevel=2) + observed_times = observed_times.astype(float) + uncensored = uncensored.astype(int) if np.sum(uncensored) == 0: raise ValueError("At least one data point must be uncensored.") if weights is None: weights = np.ones_like(observed_times) else: - assert len(weights.shape) == 1 - assert observed_times.shape[0] == weights.shape[0] + if len(weights.shape) != 1: + raise ValueError("weights must be one-dimensional array.") + if observed_times.shape[0] != weights.shape[0]: + raise ValueError("observed_times and weights must have the same length.") + if not np.issubdtype(weights.dtype, np.floating): + warnings.warn("weights is not float, converting to float.", stacklevel=2) + weights = weights.astype(float) # sort based on uncensored and observed_times temp = np.concatenate( diff --git a/src/cenreg/pytorch/cjd2F.py b/src/cenreg/pytorch/cjd2F.py index b487fc3..da33076 100644 --- a/src/cenreg/pytorch/cjd2F.py +++ b/src/cenreg/pytorch/cjd2F.py @@ -21,8 +21,10 @@ def __init__( ): super().__init__() if init_f is not None: - assert len(init_f.shape) == 3 - assert len(jd_pred.shape) == 3 + if len(init_f.shape) != 3: + raise ValueError("init_f must be a 3D array") + if len(jd_pred.shape) != 3: + raise ValueError("jd_pred must be a 3D array") self.jd_pred = torch.tensor(jd_pred, dtype=torch.float32).detach() self.focal_risk = focal_risk @@ -188,7 +190,8 @@ def minimize_mse(model, num_epochs: int) -> np.ndarray: F_pred: estimated CDF. np.ndarray of shape [batch_size, num_risks, num_bin_predictions+1] """ - assert num_epochs > 0 + if num_epochs <= 0: + raise ValueError("num_epochs must be greater than 0") best_epoch = -1 best_loss = float("inf") @@ -237,7 +240,8 @@ def minimize_mse(model, num_epochs: int) -> np.ndarray: loss.backward() optimizer.step() - assert path is not None + if path is None: + raise ValueError("path must be set to a valid checkpoint path") checkpoint = torch.load(path) model.load_state_dict(checkpoint["model_state_dict"]) model.eval() diff --git a/src/cenreg/pytorch/distribution.py b/src/cenreg/pytorch/distribution.py index 6e1fa67..5c108d5 100644 --- a/src/cenreg/pytorch/distribution.py +++ b/src/cenreg/pytorch/distribution.py @@ -404,7 +404,8 @@ def _linear_interpolation(kx: torch.Tensor, ky: torch.Tensor, x: torch.Tensor) - if kx.ndim == 1: if ky.ndim == 2: - assert kx.shape[0] == ky.shape[1] + if kx.shape[0] != ky.shape[1]: + raise ValueError("kx.shape[0] != ky.shape[1]") # compute idx and ratio if kx.ndim == 1: diff --git a/src/cenreg/pytorch/loss_cdf.py b/src/cenreg/pytorch/loss_cdf.py index be241bc..87d8beb 100644 --- a/src/cenreg/pytorch/loss_cdf.py +++ b/src/cenreg/pytorch/loss_cdf.py @@ -216,11 +216,13 @@ def brier( loss : Tensor of shape [batch_size] """ - assert len(y.shape) == 1 + if len(y.shape) != 1: + raise ValueError("y should be a 1D tensor") if y_bins is None: - y_bins = dist.boundaries - assert len(y_bins.shape) == 1 + y_bins = dist.b + if len(y_bins.shape) != 1: + raise ValueError("y_bins should be a 1D tensor") idx = torch.searchsorted(y_bins, y.view(-1, 1), right=True) F_pred = dist.cdf(y_bins) @@ -248,8 +250,10 @@ def loss( pred: torch.Tensor, y: torch.Tensor, ) -> torch.Tensor: - assert len(pred.shape) == 2 - assert len(y.shape) == 1 + if len(pred.shape) != 2: + raise ValueError("pred must be of shape [batch_size, num_bins]") + if len(y.shape) != 1: + raise ValueError("y must be of shape [batch_size]") self.distribution.set_knot_values(pred, apply_cumsum=self.apply_cumsum) return brier(self.distribution, y, self.y_bins) @@ -276,11 +280,13 @@ def ranked_probability_score( loss : Tensor of shape [batch_size] """ - assert len(y.shape) == 1 + if len(y.shape) != 1: + raise ValueError("y should be a 1D tensor") if y_bins is None: - y_bins = dist.boundaries - assert len(y_bins.shape) == 1 + y_bins = dist.b + if len(y_bins.shape) != 1: + raise ValueError("y_bins should be a 1D tensor") F_pred = dist.cdf(y_bins[1:-1]) idx = torch.searchsorted(y_bins, y.view(-1, 1), right=True) - 1 @@ -308,8 +314,10 @@ def loss( pred: torch.Tensor, y: torch.Tensor, ) -> torch.Tensor: - assert len(pred.shape) == 2 - assert len(y.shape) == 1 + if len(pred.shape) != 2: + raise ValueError("pred must be of shape [batch_size, num_bins]") + if len(y.shape) != 1: + raise ValueError("y must be of shape [batch_size]") self.distribution.set_knot_values(pred, apply_cumsum=self.apply_cumsum) return ranked_probability_score(self.distribution, y, self.y_bins) diff --git a/src/cenreg/pytorch/loss_cjd.py b/src/cenreg/pytorch/loss_cjd.py index d96a98e..452386e 100644 --- a/src/cenreg/pytorch/loss_cjd.py +++ b/src/cenreg/pytorch/loss_cjd.py @@ -14,11 +14,16 @@ def loss( observed_times: torch.Tensor, events: torch.Tensor, ) -> torch.Tensor: - assert len(pred.shape) == 2 - assert len(observed_times.shape) == 1 - assert len(events.shape) == 1 - assert observed_times.max() < self.y_bins[-1], "Observed times exceed y_bins range." - assert observed_times.min() >= self.y_bins[0], "Observed times below y_bins range." + if len(pred.shape) != 2: + raise ValueError("pred must be of shape [batch_size, num_bins]") + if len(observed_times.shape) != 1: + raise ValueError("observed_times must be of shape [batch_size]") + if len(events.shape) != 1: + raise ValueError("events must be of shape [batch_size]") + if observed_times.max() >= self.y_bins[-1]: + raise ValueError("Observed times exceed y_bins range.") + if observed_times.min() < self.y_bins[0]: + raise ValueError("Observed times below y_bins range.") events = events.long().view(-1, 1) idx = torch.searchsorted(self.y_bins, observed_times.view(-1, 1), right=True) @@ -39,11 +44,16 @@ def loss( observed_times: torch.Tensor, events: torch.Tensor, ) -> torch.Tensor: - assert len(pred.shape) == 2 - assert len(observed_times.shape) == 1 - assert len(events.shape) == 1 - assert observed_times.max() < self.y_bins[-1], "Observed times exceed y_bins range." - assert observed_times.min() >= self.y_bins[0], "Observed times below y_bins range." + if len(pred.shape) != 2: + raise ValueError("pred must be of shape [batch_size, num_bins]") + if len(observed_times.shape) != 1: + raise ValueError("observed_times must be of shape [batch_size]") + if len(events.shape) != 1: + raise ValueError("events must be of shape [batch_size]") + if observed_times.max() >= self.y_bins[-1]: + raise ValueError("Observed times exceed y_bins range.") + if observed_times.min() < self.y_bins[0]: + raise ValueError("Observed times below y_bins range.") events = events.long().view(-1, 1) idx = torch.searchsorted(self.y_bins, observed_times.view(-1, 1), right=True) diff --git a/src/cenreg/pytorch/loss_cont.py b/src/cenreg/pytorch/loss_cont.py index 126dadc..351b1fa 100644 --- a/src/cenreg/pytorch/loss_cont.py +++ b/src/cenreg/pytorch/loss_cont.py @@ -108,8 +108,13 @@ def __init__(self, copula=None, survival_copula=None, eps=0.0001): print("Warning: survival_copula is not None. copula is ignored.") def loss(self, F_pred: torch.Tensor, observed_times: torch.Tensor, events: torch.Tensor) -> torch.Tensor: - assert len(F_pred.shape) == 2 - assert F_pred.shape[0] == observed_times.shape[0] + if len(F_pred.shape) != 2: + raise ValueError("F_pred must be of shape [batch_size, num_risks]") + if F_pred.shape[0] != observed_times.shape[0]: + raise ValueError( + f"F_pred and observed_times must have the same number of samples, " + f"got {F_pred.shape[0]} and {observed_times.shape[0]}" + ) df = torch.zeros_like(observed_times) num_risks = F_pred.shape[1] diff --git a/src/cenreg/pytorch/mlp.py b/src/cenreg/pytorch/mlp.py index 2b60fb3..f84684f 100644 --- a/src/cenreg/pytorch/mlp.py +++ b/src/cenreg/pytorch/mlp.py @@ -141,8 +141,10 @@ def __init__(self, input_len: int, input_monotone_len: int, output_num: int, num self.list_smm = nn.ModuleList(list_smm) def forward(self, x: torch.Tensor, t: torch.Tensor) -> torch.Tensor: - assert len(x.shape) == 2 - assert len(t.shape) == 2 + if len(x.shape) != 2: + raise ValueError("x must be of shape [batch_size, num_features]") + if len(t.shape) != 2: + raise ValueError("t must be of shape [batch_size, num_monotone_features]") x = F.relu(self.fc1(x)) list_out = [] diff --git a/src/cenreg/pytorch/utils.py b/src/cenreg/pytorch/utils.py index 86342af..6d547ab 100644 --- a/src/cenreg/pytorch/utils.py +++ b/src/cenreg/pytorch/utils.py @@ -20,7 +20,8 @@ def normalize_y(y: torch.Tensor, min_y: float, max_y: float) -> torch.Tensor: The normalized tensor. """ - assert min_y < max_y + if min_y >= max_y: + raise ValueError("min_y must be strictly less than max_y") return (y - min_y) / (max_y - min_y) @@ -44,10 +45,14 @@ def denormalize_pred(pred: torch.Tensor, min_y: float, max_y: float) -> torch.Te The denormalized predictions. """ - assert pred.dim() == 2 - assert pred.shape[1] > 0 - assert min_y < max_y - assert (pred.min() >= 0) and (pred.max() <= 1) + if pred.dim() != 2: + raise ValueError("pred must be of shape [batch_size, num_features]") + if pred.shape[1] <= 0: + raise ValueError("pred must have at least one feature") + if min_y >= max_y: + raise ValueError("min_y must be strictly less than max_y") + if (pred.min() < 0) or (pred.max() > 1): + raise ValueError("pred values must be in the range [0, 1]") pred_cumsum = torch.cumsum(pred, dim=1) pred_cumsum = torch.cat([torch.zeros(pred.shape[0], 1, device=pred.device), pred_cumsum], dim=1) diff --git a/src/cenreg/utils.py b/src/cenreg/utils.py index b280182..55c320a 100644 --- a/src/cenreg/utils.py +++ b/src/cenreg/utils.py @@ -21,8 +21,10 @@ def create_bins(max_y: float, min_y: float = 0.0, num_bins=10, algorithm: str = bins : np.ndarray Array of bin edges. """ - assert num_bins > 1, "Number of bins must be greater than 1." - assert max_y > min_y, "Maximum value must be greater than minimum value." + if num_bins <= 1: + raise ValueError("Number of bins must be greater than 1.") + if max_y <= min_y: + raise ValueError("Maximum value must be strictly greater than minimum value.") if algorithm != "even": raise ValueError(f"Unknown algorithm: {algorithm}. Supported: 'even'.")