diff --git a/src/cenreg/distribution/cdf.py b/src/cenreg/distribution/cdf.py index db888e1..ee43381 100644 --- a/src/cenreg/distribution/cdf.py +++ b/src/cenreg/distribution/cdf.py @@ -1,3 +1,5 @@ +from typing import Literal + import numpy as np from cenreg.distribution.interpolate import linear_interpolation @@ -13,7 +15,7 @@ def __init__( b: np.ndarray, p: np.ndarray | None = None, cum_p: np.ndarray | None = None, - interpolate: str = "linear", + interpolate: Literal["linear", "left", "right"] = "linear", confidence_interval: np.ndarray | None = None, ): """ @@ -31,7 +33,7 @@ def __init__( Cumulative probability distribution. cum_p must be one-dimensional or two-dimensional. If both p and cum_p are given, cum_p is used. - interpolate : str + interpolate : Literal["linear", "left", "right"] 'linear', 'left', or 'right' indicating the interpolation method. If 'linear' is set, linear interpolation is used. If 'left' is set, the CDF value at the left edge of each bin is used. @@ -73,7 +75,7 @@ def __init__( self.interpolate = interpolate self.confidence_interval = confidence_interval - def cdf(self, y: float | np.ndarray): + def cdf(self, y: int | float | np.ndarray): """ Cumulative distribution function (i.e., inverse of quantile function). @@ -87,6 +89,8 @@ def cdf(self, y: float | np.ndarray): cum_p : np.ndarray CDF values for each value in y. """ + if isinstance(y, int | float): + y = np.array([y], dtype=float) if self.cum_p.ndim == 1: assert y.ndim == 1 @@ -95,8 +99,6 @@ def cdf(self, y: float | np.ndarray): else: raise ValueError("cum_p must be one-dimensional or two-dimensional.") - if isinstance(y, float): - y = np.array([y]) if self.cum_p.ndim == 2 and y.ndim == 1: y = np.tile(y, (self.cum_p.shape[0], 1)) @@ -154,8 +156,8 @@ def icdf(self, quantiles: float | np.ndarray) -> np.ndarray: Compute inverse CDF values for each value in quantiles. """ - if isinstance(quantiles, float): - quantiles = np.array([quantiles]) + if isinstance(quantiles, int | float): + quantiles = np.array([quantiles], dtype=float) if np.any(quantiles < 0.0): raise ValueError("quantiles must be non-negative.") if np.any(quantiles > 1.0): diff --git a/src/cenreg/distribution/quantile.py b/src/cenreg/distribution/quantile.py index b2101f7..e053433 100644 --- a/src/cenreg/distribution/quantile.py +++ b/src/cenreg/distribution/quantile.py @@ -61,8 +61,8 @@ def cdf(self, y: float | np.ndarray): CDF values for each value in y. Array shape is equal to the shape of y. """ - if isinstance(y, float): - y = np.array([y]) + if isinstance(y, int | float): + y = np.array([y], dtype=float) if self.interpolate == "linear": # linear interpolation implementation @@ -97,7 +97,7 @@ def icdf(self, quantiles: float | np.ndarray) -> np.ndarray: Array shape is equal to the shape of quantiles. """ - if isinstance(quantiles, float): + if isinstance(quantiles, int | float): quantiles = np.array([quantiles]) if np.any(quantiles < 0.0): raise ValueError("quantiles must be non-negative.") diff --git a/src/cenreg/model/nonparametric.py b/src/cenreg/model/nonparametric.py index 0f3b349..58767e2 100644 --- a/src/cenreg/model/nonparametric.py +++ b/src/cenreg/model/nonparametric.py @@ -205,8 +205,8 @@ def kaplan_meier_estimator( else: survival_rates = survival_rates[:-1] dist = CumulativeDist(b=b, cum_p=1.0 - survival_rates, interpolate="right") - dist.alive = num_alive - dist.dead = num_death + # dist.alive = num_alive + # dist.dead = num_death return dist diff --git a/src/cenreg/pytorch/cjd2F.py b/src/cenreg/pytorch/cjd2F.py index 2a4e0e1..b487fc3 100644 --- a/src/cenreg/pytorch/cjd2F.py +++ b/src/cenreg/pytorch/cjd2F.py @@ -20,13 +20,14 @@ def __init__( optimizer=None, ): super().__init__() - assert len(init_f.shape) == 3 + if init_f is not None: + assert len(init_f.shape) == 3 assert len(jd_pred.shape) == 3 self.jd_pred = torch.tensor(jd_pred, dtype=torch.float32).detach() self.focal_risk = focal_risk self.fc = nn.Linear(1, jd_pred.size, bias=False) - self.shape = init_f.shape + self.shape = init_f.shape if init_f is not None else jd_pred.shape self.copula = copula self.learning_rate = learning_rate if init_f is not None: @@ -61,7 +62,7 @@ def _copula_sum_sub( self, F_pred: torch.Tensor, c, - idx_list: list[int], + idx_list: list[list[int]], i: int, k: int, idx_list_use_Ft: Sequence[int], @@ -187,6 +188,7 @@ 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 best_epoch = -1 best_loss = float("inf") @@ -235,6 +237,7 @@ def minimize_mse(model, num_epochs: int) -> np.ndarray: loss.backward() optimizer.step() + assert path is not None checkpoint = torch.load(path) model.load_state_dict(checkpoint["model_state_dict"]) model.eval()