Skip to content

Analysis

Summaries, quality checks and plots for a generated dataset.

Task-oriented walkthroughs: Summarising a dataset and Plotting.

Summaries

summary

Utilities for summarizing and validating survival datasets.

This module provides functions to summarize survival data, check data quality, and identify potential issues.

summarize_survival_dataset

summarize_survival_dataset(
    data: DataFrame,
    time_col: str = "time",
    status_col: str = "status",
    id_col: str | None = None,
    covariate_cols: list[str] | None = None,
    verbose: bool = True,
) -> dict[str, Any]

Generate a comprehensive summary of a survival dataset.

Parameters:

Name Type Description Default
data DataFrame

DataFrame containing survival data.

required
time_col str

Name of the column containing time-to-event values.

"time"
status_col str

Name of the column containing event indicators (1=event, 0=censored).

"status"
id_col str

Name of the column containing subject identifiers.

None
covariate_cols list of str

List of column names to include as covariates in the summary. If None, all columns except time_col, status_col, and id_col are considered.

None
verbose bool

Whether to print the summary to console.

True

Returns:

Type Description
dict[str, Any]

Dictionary containing all summary statistics.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.summary import summarize_survival_dataset
>>>
>>> # Generate example data
>>> df = generate(model="cphm", n=100, model_cens="uniform",
...               cens_par=1.0, beta=0.5, covariate_range=2.0)
>>>
>>> # Summarize the dataset
>>> summary = summarize_survival_dataset(df)
Source code in gen_surv/summary.py
def summarize_survival_dataset(
    data: pd.DataFrame,
    time_col: str = "time",
    status_col: str = "status",
    id_col: str | None = None,
    covariate_cols: list[str] | None = None,
    verbose: bool = True,
) -> dict[str, Any]:
    """
    Generate a comprehensive summary of a survival dataset.

    Parameters
    ----------
    data : pd.DataFrame
        DataFrame containing survival data.
    time_col : str, default="time"
        Name of the column containing time-to-event values.
    status_col : str, default="status"
        Name of the column containing event indicators (1=event, 0=censored).
    id_col : str, optional
        Name of the column containing subject identifiers.
    covariate_cols : list of str, optional
        List of column names to include as covariates in the summary.
        If None, all columns except time_col, status_col, and id_col are considered.
    verbose : bool, default=True
        Whether to print the summary to console.

    Returns
    -------
    dict[str, Any]
        Dictionary containing all summary statistics.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.summary import summarize_survival_dataset
    >>>
    >>> # Generate example data
    >>> df = generate(model="cphm", n=100, model_cens="uniform",
    ...               cens_par=1.0, beta=0.5, covariate_range=2.0)
    >>>
    >>> # Summarize the dataset
    >>> summary = summarize_survival_dataset(df)
    """
    # Validate input columns
    for col in [time_col, status_col]:
        if col not in data.columns:
            raise ParameterError("column", col, "not found in data")

    if id_col is not None and id_col not in data.columns:
        raise ParameterError("id_col", id_col, "not found in data")

    # Determine covariate columns
    if covariate_cols is None:
        exclude_cols = {time_col, status_col}
        if id_col is not None:
            exclude_cols.add(id_col)
        covariate_cols = [col for col in data.columns if col not in exclude_cols]
    else:
        missing_cols = [col for col in covariate_cols if col not in data.columns]
        if missing_cols:
            raise ParameterError("covariate_cols", missing_cols, "not found in data")

    # Basic dataset information
    n_subjects = len(data)
    if id_col is not None:
        n_unique_ids = data[id_col].nunique()
    else:
        n_unique_ids = n_subjects

    # Event information
    n_events = data[status_col].sum()
    n_censored = n_subjects - n_events
    event_rate = n_events / n_subjects

    # Time statistics
    time_min = data[time_col].min()
    time_max = data[time_col].max()
    time_mean = data[time_col].mean()
    time_median = data[time_col].median()

    # Data quality checks
    n_missing_time = data[time_col].isna().sum()
    n_missing_status = data[status_col].isna().sum()
    n_negative_time = (data[time_col] < 0).sum()
    n_invalid_status = data[~data[status_col].isin([0, 1])].shape[0]

    # Covariate summaries
    covariate_stats = {}
    for col in covariate_cols:
        col_data = data[col]
        is_numeric = pd.api.types.is_numeric_dtype(col_data)

        if is_numeric:
            covariate_stats[col] = {
                "type": "numeric",
                "min": col_data.min(),
                "max": col_data.max(),
                "mean": col_data.mean(),
                "median": col_data.median(),
                "std": col_data.std(),
                "missing": col_data.isna().sum(),
                "unique_values": col_data.nunique(),
            }
        else:
            # Categorical/string
            covariate_stats[col] = {
                "type": "categorical",
                "n_categories": col_data.nunique(),
                "top_categories": col_data.value_counts().head(5).to_dict(),
                "missing": col_data.isna().sum(),
            }

    # Compile the summary
    summary = {
        "dataset_info": {
            "n_subjects": n_subjects,
            "n_unique_ids": n_unique_ids,
            "n_covariates": len(covariate_cols),
        },
        "event_info": {
            "n_events": n_events,
            "n_censored": n_censored,
            "event_rate": event_rate,
        },
        "time_info": {
            "min": time_min,
            "max": time_max,
            "mean": time_mean,
            "median": time_median,
        },
        "data_quality": {
            "missing_time": n_missing_time,
            "missing_status": n_missing_status,
            "negative_time": n_negative_time,
            "invalid_status": n_invalid_status,
            "overall_quality": (
                "good"
                if (
                    n_missing_time
                    + n_missing_status
                    + n_negative_time
                    + n_invalid_status
                )
                == 0
                else "issues_detected"
            ),
        },
        "covariates": covariate_stats,
    }

    # Print summary if requested
    if verbose:
        _print_summary(summary, time_col, status_col, id_col, covariate_cols)

    return summary

check_survival_data_quality

check_survival_data_quality(
    data: DataFrame,
    time_col: str = "time",
    status_col: str = "status",
    id_col: str | None = None,
    min_time: float = 0.0,
    max_time: float | None = None,
    status_values: list[int] | None = None,
    fix_issues: bool = False,
) -> tuple[DataFrame, dict[str, Any]]

Check for common issues in survival data and optionally fix them.

Parameters:

Name Type Description Default
data DataFrame

DataFrame containing survival data.

required
time_col str

Name of the column containing time-to-event values.

"time"
status_col str

Name of the column containing event indicators.

"status"
id_col str

Name of the column containing subject identifiers.

None
min_time float

Minimum acceptable value for time column.

0.0
max_time float

Maximum acceptable value for time column.

None
status_values list of int

List of valid status values. Default is [0, 1].

None
fix_issues bool

Whether to attempt fixing issues (returns a modified DataFrame).

False

Returns:

Type Description
tuple[DataFrame, dict[str, Any]]

Tuple containing (possibly fixed) DataFrame and issues report.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.summary import check_survival_data_quality
>>>
>>> # Generate example data with some issues
>>> df = generate(model="cphm", n=100, model_cens="uniform",
...               cens_par=1.0, beta=0.5, covariate_range=2.0)
>>> # Introduce some issues
>>> df.loc[0, "time"] = np.nan
>>> df.loc[1, "status"] = 2  # Invalid status
>>>
>>> # Check and fix issues
>>> fixed_df, issues = check_survival_data_quality(df, fix_issues=True)
>>> print(issues)
Source code in gen_surv/summary.py
def check_survival_data_quality(
    data: pd.DataFrame,
    time_col: str = "time",
    status_col: str = "status",
    id_col: str | None = None,
    min_time: float = 0.0,
    max_time: float | None = None,
    status_values: list[int] | None = None,
    fix_issues: bool = False,
) -> tuple[pd.DataFrame, dict[str, Any]]:
    """
    Check for common issues in survival data and optionally fix them.

    Parameters
    ----------
    data : pd.DataFrame
        DataFrame containing survival data.
    time_col : str, default="time"
        Name of the column containing time-to-event values.
    status_col : str, default="status"
        Name of the column containing event indicators.
    id_col : str, optional
        Name of the column containing subject identifiers.
    min_time : float, default=0.0
        Minimum acceptable value for time column.
    max_time : float, optional
        Maximum acceptable value for time column.
    status_values : list of int, optional
        List of valid status values. Default is [0, 1].
    fix_issues : bool, default=False
        Whether to attempt fixing issues (returns a modified DataFrame).

    Returns
    -------
    tuple[pd.DataFrame, dict[str, Any]]
        Tuple containing (possibly fixed) DataFrame and issues report.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.summary import check_survival_data_quality
    >>>
    >>> # Generate example data with some issues
    >>> df = generate(model="cphm", n=100, model_cens="uniform",
    ...               cens_par=1.0, beta=0.5, covariate_range=2.0)
    >>> # Introduce some issues
    >>> df.loc[0, "time"] = np.nan
    >>> df.loc[1, "status"] = 2  # Invalid status
    >>>
    >>> # Check and fix issues
    >>> fixed_df, issues = check_survival_data_quality(df, fix_issues=True)
    >>> print(issues)
    """
    if status_values is None:
        status_values = [0, 1]

    # Make a copy to avoid modifying the original
    if fix_issues:
        data = data.copy()

    # Initialize issues report
    issues: dict[str, dict[str, int | None]] = {
        "missing_data": {"time": 0, "status": 0, "id": 0 if id_col else None},
        "invalid_values": {
            "negative_time": 0,
            "excessive_time": 0,
            "invalid_status": 0,
        },
        "duplicates": {"duplicate_rows": 0, "duplicate_ids": 0 if id_col else None},
        "modifications": {"rows_dropped": 0, "values_fixed": 0},
    }

    # Check for missing values
    issues["missing_data"]["time"] = data[time_col].isna().sum()
    issues["missing_data"]["status"] = data[status_col].isna().sum()
    if id_col:
        issues["missing_data"]["id"] = data[id_col].isna().sum()

    # Check for invalid values
    issues["invalid_values"]["negative_time"] = (data[time_col] < min_time).sum()
    if max_time is not None:
        issues["invalid_values"]["excessive_time"] = (data[time_col] > max_time).sum()
    issues["invalid_values"]["invalid_status"] = data[
        ~data[status_col].isin(status_values)
    ].shape[0]

    # Check for duplicates
    issues["duplicates"]["duplicate_rows"] = data.duplicated().sum()
    if id_col:
        issues["duplicates"]["duplicate_ids"] = data[id_col].duplicated().sum()

    # Fix issues if requested
    if fix_issues:
        original_rows = len(data)
        modified_values = 0

        # Handle missing values
        data = data.dropna(subset=[time_col, status_col])

        # Handle invalid values
        if min_time > 0:
            # Set negative or too small times to min_time
            mask = data[time_col] < min_time
            if mask.any():
                data.loc[mask, time_col] = min_time
                modified_values += mask.sum()

        if max_time is not None:
            # Cap excessively large times
            mask = data[time_col] > max_time
            if mask.any():
                data.loc[mask, time_col] = max_time
                modified_values += mask.sum()

        # Fix invalid status values
        mask = ~data[status_col].isin(status_values)
        if mask.any():
            # Default to censored (0) for invalid status
            data.loc[mask, status_col] = 0
            modified_values += mask.sum()

        # Remove duplicates
        data = data.drop_duplicates()

        # Update modification counts
        issues["modifications"]["rows_dropped"] = original_rows - len(data)
        issues["modifications"]["values_fixed"] = modified_values

    return data, issues

compare_survival_datasets

compare_survival_datasets(
    datasets: dict[str, DataFrame],
    time_col: str = "time",
    status_col: str = "status",
    covariate_cols: list[str] | None = None,
) -> DataFrame

Compare multiple survival datasets and summarize their differences.

Parameters:

Name Type Description Default
datasets dict[str, DataFrame]

Dictionary mapping dataset names to DataFrames.

required
time_col str

Name of the time column in each dataset.

"time"
status_col str

Name of the status column in each dataset.

"status"
covariate_cols List[str]

List of covariate columns to compare. If None, compares all common columns.

None

Returns:

Type Description
DataFrame

Comparison table with datasets as columns and metrics as rows.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.summary import compare_survival_datasets
>>>
>>> # Generate datasets with different parameters
>>> datasets = {
...     "CPHM": generate(model="cphm", n=100, model_cens="uniform",
...                    cens_par=1.0, beta=0.5, covariate_range=2.0),
...     "Weibull AFT": generate(model="aft_weibull", n=100, beta=[0.5],
...                           shape=1.5, scale=1.0, model_cens="uniform", cens_par=1.0)
... }
>>>
>>> # Compare datasets
>>> comparison = compare_survival_datasets(datasets)
>>> print(comparison)
Source code in gen_surv/summary.py
def compare_survival_datasets(
    datasets: dict[str, pd.DataFrame],
    time_col: str = "time",
    status_col: str = "status",
    covariate_cols: list[str] | None = None,
) -> pd.DataFrame:
    """
    Compare multiple survival datasets and summarize their differences.

    Parameters
    ----------
    datasets : dict[str, pd.DataFrame]
        Dictionary mapping dataset names to DataFrames.
    time_col : str, default="time"
        Name of the time column in each dataset.
    status_col : str, default="status"
        Name of the status column in each dataset.
    covariate_cols : List[str], optional
        List of covariate columns to compare. If None, compares all common columns.

    Returns
    -------
    pd.DataFrame
        Comparison table with datasets as columns and metrics as rows.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.summary import compare_survival_datasets
    >>>
    >>> # Generate datasets with different parameters
    >>> datasets = {
    ...     "CPHM": generate(model="cphm", n=100, model_cens="uniform",
    ...                    cens_par=1.0, beta=0.5, covariate_range=2.0),
    ...     "Weibull AFT": generate(model="aft_weibull", n=100, beta=[0.5],
    ...                           shape=1.5, scale=1.0, model_cens="uniform", cens_par=1.0)
    ... }
    >>>
    >>> # Compare datasets
    >>> comparison = compare_survival_datasets(datasets)
    >>> print(comparison)
    """
    if not datasets:
        raise ParameterError("datasets", datasets, "at least one dataset is required")

    # Find common columns if covariate_cols not specified
    if covariate_cols is None:
        all_columns = [set(df.columns) for df in datasets.values()]
        common_columns = set.intersection(*all_columns)
        common_columns -= {time_col, status_col}  # Remove time and status
        covariate_cols = sorted(list(common_columns))

    # Calculate summaries for each dataset
    summaries = {}
    for name, data in datasets.items():
        summaries[name] = summarize_survival_dataset(
            data, time_col, status_col, covariate_cols=covariate_cols, verbose=False
        )

    # Construct the comparison DataFrame
    comparison_data = {}

    # Dataset info
    comparison_data["n_subjects"] = {
        name: summary["dataset_info"]["n_subjects"]
        for name, summary in summaries.items()
    }
    comparison_data["n_events"] = {
        name: summary["event_info"]["n_events"] for name, summary in summaries.items()
    }
    comparison_data["event_rate"] = {
        name: summary["event_info"]["event_rate"] for name, summary in summaries.items()
    }

    # Time info
    comparison_data["time_min"] = {
        name: summary["time_info"]["min"] for name, summary in summaries.items()
    }
    comparison_data["time_max"] = {
        name: summary["time_info"]["max"] for name, summary in summaries.items()
    }
    comparison_data["time_mean"] = {
        name: summary["time_info"]["mean"] for name, summary in summaries.items()
    }
    comparison_data["time_median"] = {
        name: summary["time_info"]["median"] for name, summary in summaries.items()
    }

    # Covariate info (means for numeric)
    for col in covariate_cols:
        for name, summary in summaries.items():
            if col in summary["covariates"]:
                col_stats = summary["covariates"][col]
                if col_stats["type"] == "numeric":
                    if f"{col}_mean" not in comparison_data:
                        comparison_data[f"{col}_mean"] = {}
                    comparison_data[f"{col}_mean"][name] = col_stats["mean"]

    # Create the DataFrame
    comparison_df = pd.DataFrame(comparison_data).T

    return comparison_df

Plots

visualization

Visualization utilities for survival data.

This module provides functions to visualize survival data generated by gen_surv, including Kaplan-Meier survival curves and other commonly used plots in survival analysis.

plot_survival_curve

plot_survival_curve(
    data: DataFrame,
    time_col: str = "time",
    status_col: str = "status",
    group_col: str | None = None,
    confidence_intervals: bool = True,
    title: str = "Kaplan-Meier Survival Curve",
    figsize: tuple[float, float] = (10, 6),
    ci_alpha: float = 0.2,
) -> tuple[Figure, Axes]

Plot Kaplan-Meier survival curves from simulated data.

Parameters:

Name Type Description Default
data DataFrame

DataFrame containing the survival data.

required
time_col str

Name of the column containing event/censoring times.

"time"
status_col str

Name of the column containing event indicators (1=event, 0=censored).

"status"
group_col str

Name of the column to use for stratification (creates separate curves).

None
confidence_intervals bool

Whether to display confidence intervals around the survival curves.

True
title str

Plot title.

"Kaplan-Meier Survival Curve"
figsize tuple

Figure size (width, height) in inches.

(10, 6)
ci_alpha float

Transparency level for confidence interval bands.

0.2

Returns:

Name Type Description
fig Figure

Matplotlib figure object.

ax Axes

Matplotlib axes object.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.visualization import plot_survival_curve
>>>
>>> # Generate data
>>> df = generate(model="cphm", n=100, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0)
>>>
>>> # Create a categorical group based on covariate
>>> df["group"] = pd.cut(df["covariate"], bins=2, labels=["Low", "High"])
>>>
>>> # Plot survival curves by group
>>> fig, ax = plot_survival_curve(df, group_col="group")
>>> plt.show()
Source code in gen_surv/visualization.py
def plot_survival_curve(
    data: pd.DataFrame,
    time_col: str = "time",
    status_col: str = "status",
    group_col: str | None = None,
    confidence_intervals: bool = True,
    title: str = "Kaplan-Meier Survival Curve",
    figsize: tuple[float, float] = (10, 6),
    ci_alpha: float = 0.2,
) -> tuple[Figure, Axes]:
    """
    Plot Kaplan-Meier survival curves from simulated data.

    Parameters
    ----------
    data : pd.DataFrame
        DataFrame containing the survival data.
    time_col : str, default="time"
        Name of the column containing event/censoring times.
    status_col : str, default="status"
        Name of the column containing event indicators (1=event, 0=censored).
    group_col : str, optional
        Name of the column to use for stratification (creates separate curves).
    confidence_intervals : bool, default=True
        Whether to display confidence intervals around the survival curves.
    title : str, default="Kaplan-Meier Survival Curve"
        Plot title.
    figsize : tuple, default=(10, 6)
        Figure size (width, height) in inches.
    ci_alpha : float, default=0.2
        Transparency level for confidence interval bands.

    Returns
    -------
    fig : Figure
        Matplotlib figure object.
    ax : Axes
        Matplotlib axes object.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.visualization import plot_survival_curve
    >>>
    >>> # Generate data
    >>> df = generate(model="cphm", n=100, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0)
    >>>
    >>> # Create a categorical group based on covariate
    >>> df["group"] = pd.cut(df["covariate"], bins=2, labels=["Low", "High"])
    >>>
    >>> # Plot survival curves by group
    >>> fig, ax = plot_survival_curve(df, group_col="group")
    >>> plt.show()
    """
    # Import lifelines here to avoid making it a hard dependency
    try:
        from lifelines import KaplanMeierFitter
        from lifelines.plotting import add_at_risk_counts
    except ImportError as exc:
        raise ImportError(
            "This function requires the lifelines package. "
            "Install it with: pip install lifelines"
        ) from exc

    fig, ax = plt.subplots(figsize=figsize)

    # Create separate KM curves for each group (if specified)
    if group_col is not None:
        groups = data[group_col].unique()
        cmap = plt.get_cmap("tab10")
        colors = [cmap(i) for i in range(len(groups))]

        for i, group in enumerate(groups):
            mask = data[group_col] == group
            group_data = data[mask]

            kmf = KaplanMeierFitter()
            kmf.fit(
                group_data[time_col],
                group_data[status_col],
                label=f"{group_col}={group}",
            )

            kmf.plot_survival_function(
                ax=ax, ci_show=confidence_intervals, color=colors[i], ci_alpha=ci_alpha
            )

        # Add at-risk counts below the plot
        add_at_risk_counts(kmf, ax=ax)
    else:
        # Single KM curve for all data
        kmf = KaplanMeierFitter()
        kmf.fit(data[time_col], data[status_col])

        kmf.plot_survival_function(
            ax=ax, ci_show=confidence_intervals, ci_alpha=ci_alpha
        )

        # Add at-risk counts below the plot
        add_at_risk_counts(kmf, ax=ax)

    # Customize plot appearance
    ax.set_title(title)
    ax.set_xlabel("Time")
    ax.set_ylabel("Survival Probability")
    ax.grid(alpha=0.3)
    ax.set_ylim(0, 1.05)

    plt.tight_layout()
    return fig, ax

plot_hazard_comparison

plot_hazard_comparison(
    models: dict[str, DataFrame],
    time_col: str = "time",
    status_col: str = "status",
    title: str = "Hazard Function Comparison",
    figsize: tuple[float, float] = (10, 6),
    bandwidth: float = 0.5,
) -> tuple[Figure, Axes]

Compare hazard functions from multiple generated datasets.

Parameters:

Name Type Description Default
models dict

Dictionary mapping model names to their respective DataFrames.

required
time_col str

Name of the column containing event/censoring times.

"time"
status_col str

Name of the column containing event indicators (1=event, 0=censored).

"status"
title str

Plot title.

"Hazard Function Comparison"
figsize tuple

Figure size (width, height) in inches.

(10, 6)
bandwidth float

Bandwidth parameter for kernel density estimation of the hazard function.

0.5

Returns:

Name Type Description
fig Figure

Matplotlib figure object.

ax Axes

Matplotlib axes object.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.visualization import plot_hazard_comparison
>>>
>>> # Generate data from multiple models
>>> models = {
>>>     "CPHM": generate(model="cphm", n=100, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0),
>>>     "AFT Weibull": generate(model="aft_weibull", n=100, beta=[0.5], shape=1.5, scale=2.0,
>>>                            model_cens="uniform", cens_par=1.0)
>>> }
>>>
>>> # Compare hazard functions
>>> fig, ax = plot_hazard_comparison(models)
>>> plt.show()
Source code in gen_surv/visualization.py
def plot_hazard_comparison(
    models: dict[str, pd.DataFrame],
    time_col: str = "time",
    status_col: str = "status",
    title: str = "Hazard Function Comparison",
    figsize: tuple[float, float] = (10, 6),
    bandwidth: float = 0.5,
) -> tuple[Figure, Axes]:
    """
    Compare hazard functions from multiple generated datasets.

    Parameters
    ----------
    models : dict
        Dictionary mapping model names to their respective DataFrames.
    time_col : str, default="time"
        Name of the column containing event/censoring times.
    status_col : str, default="status"
        Name of the column containing event indicators (1=event, 0=censored).
    title : str, default="Hazard Function Comparison"
        Plot title.
    figsize : tuple, default=(10, 6)
        Figure size (width, height) in inches.
    bandwidth : float, default=0.5
        Bandwidth parameter for kernel density estimation of the hazard function.

    Returns
    -------
    fig : Figure
        Matplotlib figure object.
    ax : Axes
        Matplotlib axes object.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.visualization import plot_hazard_comparison
    >>>
    >>> # Generate data from multiple models
    >>> models = {
    >>>     "CPHM": generate(model="cphm", n=100, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0),
    >>>     "AFT Weibull": generate(model="aft_weibull", n=100, beta=[0.5], shape=1.5, scale=2.0,
    >>>                            model_cens="uniform", cens_par=1.0)
    >>> }
    >>>
    >>> # Compare hazard functions
    >>> fig, ax = plot_hazard_comparison(models)
    >>> plt.show()
    """
    # Import lifelines here to avoid making it a hard dependency
    try:
        from lifelines import NelsonAalenFitter
    except ImportError as exc:
        raise ImportError(
            "This function requires the lifelines package. "
            "Install it with: pip install lifelines"
        ) from exc

    fig, ax = plt.subplots(figsize=figsize)

    for model_name, df in models.items():
        naf = NelsonAalenFitter()
        naf.fit(df[time_col], df[status_col])

        # Get smoothed hazard estimate
        hazard = naf.smoothed_hazard_(bandwidth=bandwidth)

        # Plot hazard function
        ax.plot(hazard.index, hazard.values, label=model_name, alpha=0.8)

    # Customize plot appearance
    ax.set_title(title)
    ax.set_xlabel("Time")
    ax.set_ylabel("Hazard Rate")
    ax.grid(alpha=0.3)
    ax.legend()

    plt.tight_layout()
    return fig, ax

plot_covariate_effect

plot_covariate_effect(
    data: DataFrame,
    covariate_col: str,
    time_col: str = "time",
    status_col: str = "status",
    n_groups: int = 3,
    title: str = "Effect of Covariate on Survival",
    figsize: tuple[float, float] = (10, 6),
    ci_alpha: float = 0.2,
) -> tuple[Figure, Axes]

Visualize the effect of a continuous covariate on survival by discretizing it.

Parameters:

Name Type Description Default
data DataFrame

DataFrame containing the survival data.

required
covariate_col str

Name of the covariate column to visualize.

required
time_col str

Name of the column containing event/censoring times.

"time"
status_col str

Name of the column containing event indicators (1=event, 0=censored).

"status"
n_groups int

Number of groups to divide the covariate into (e.g., 3 for tertiles).

3
title str

Plot title.

"Effect of Covariate on Survival"
figsize tuple

Figure size (width, height) in inches.

(10, 6)
ci_alpha float

Transparency level for confidence interval bands.

0.2

Returns:

Name Type Description
fig Figure

Matplotlib figure object.

ax Axes

Matplotlib axes object.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.visualization import plot_covariate_effect
>>>
>>> # Generate data with a continuous covariate
>>> df = generate(model="cphm", n=200, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0)
>>>
>>> # Visualize the effect of the covariate on survival
>>> fig, ax = plot_covariate_effect(df, covariate_col="covariate", n_groups=3)
>>> plt.show()
Source code in gen_surv/visualization.py
def plot_covariate_effect(
    data: pd.DataFrame,
    covariate_col: str,
    time_col: str = "time",
    status_col: str = "status",
    n_groups: int = 3,
    title: str = "Effect of Covariate on Survival",
    figsize: tuple[float, float] = (10, 6),
    ci_alpha: float = 0.2,
) -> tuple[Figure, Axes]:
    """
    Visualize the effect of a continuous covariate on survival by discretizing
    it.

    Parameters
    ----------
    data : pd.DataFrame
        DataFrame containing the survival data.
    covariate_col : str
        Name of the covariate column to visualize.
    time_col : str, default="time"
        Name of the column containing event/censoring times.
    status_col : str, default="status"
        Name of the column containing event indicators (1=event, 0=censored).
    n_groups : int, default=3
        Number of groups to divide the covariate into (e.g., 3 for tertiles).
    title : str, default="Effect of Covariate on Survival"
        Plot title.
    figsize : tuple, default=(10, 6)
        Figure size (width, height) in inches.
    ci_alpha : float, default=0.2
        Transparency level for confidence interval bands.

    Returns
    -------
    fig : Figure
        Matplotlib figure object.
    ax : Axes
        Matplotlib axes object.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.visualization import plot_covariate_effect
    >>>
    >>> # Generate data with a continuous covariate
    >>> df = generate(model="cphm", n=200, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0)
    >>>
    >>> # Visualize the effect of the covariate on survival
    >>> fig, ax = plot_covariate_effect(df, covariate_col="covariate", n_groups=3)
    >>> plt.show()
    """
    # Add a categorical version of the covariate
    group_labels = [f"Q{i + 1}" for i in range(n_groups)]
    data = data.copy()
    data["_group"] = pd.qcut(data[covariate_col], q=n_groups, labels=group_labels)

    # Get the median value of each group for the legend
    group_medians = data.groupby("_group")[covariate_col].median()

    # Create more informative labels
    label_map = {
        group: f"{group} ({covariate_col}{median:.2f})"
        for group, median in group_medians.items()
    }

    data["_label"] = data["_group"].map(label_map)

    # Create the plot
    fig, ax = plot_survival_curve(
        data=data,
        time_col=time_col,
        status_col=status_col,
        group_col="_label",
        confidence_intervals=True,
        title=title,
        figsize=figsize,
        ci_alpha=ci_alpha,
    )

    return fig, ax

describe_survival

describe_survival(
    data: DataFrame,
    time_col: str = "time",
    status_col: str = "status",
) -> DataFrame

Generate a summary of survival data including median survival time, event counts, and other descriptive statistics.

Parameters:

Name Type Description Default
data DataFrame

DataFrame containing the survival data.

required
time_col str

Name of the column containing event/censoring times.

"time"
status_col str

Name of the column containing event indicators (1=event, 0=censored).

"status"

Returns:

Type Description
DataFrame

Summary statistics dataframe.

Examples:

>>> from gen_surv import generate
>>> from gen_surv.visualization import describe_survival
>>>
>>> # Generate data
>>> df = generate(model="cphm", n=200, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0)
>>>
>>> # Get survival summary
>>> summary = describe_survival(df)
>>> print(summary)
Source code in gen_surv/visualization.py
def describe_survival(
    data: pd.DataFrame, time_col: str = "time", status_col: str = "status"
) -> pd.DataFrame:
    """
    Generate a summary of survival data including median survival time,
    event counts, and other descriptive statistics.

    Parameters
    ----------
    data : pd.DataFrame
        DataFrame containing the survival data.
    time_col : str, default="time"
        Name of the column containing event/censoring times.
    status_col : str, default="status"
        Name of the column containing event indicators (1=event, 0=censored).

    Returns
    -------
    pd.DataFrame
        Summary statistics dataframe.

    Examples
    --------
    >>> from gen_surv import generate
    >>> from gen_surv.visualization import describe_survival
    >>>
    >>> # Generate data
    >>> df = generate(model="cphm", n=200, model_cens="uniform", cens_par=1.0, beta=0.5, covariate_range=2.0)
    >>>
    >>> # Get survival summary
    >>> summary = describe_survival(df)
    >>> print(summary)
    """
    # Import lifelines here to avoid making it a hard dependency
    try:
        from lifelines import KaplanMeierFitter
    except ImportError as exc:
        raise ImportError(
            "This function requires the lifelines package. "
            "Install it with: pip install lifelines"
        ) from exc

    n_total = len(data)
    n_events = data[status_col].sum()
    n_censored = n_total - n_events
    event_rate = n_events / n_total

    # Calculate median and other percentiles
    kmf = KaplanMeierFitter()
    kmf.fit(data[time_col], data[status_col])
    median = kmf.median_survival_time_

    # Time ranges
    time_min = data[time_col].min()
    time_max = data[time_col].max()
    time_mean = data[time_col].mean()

    # Create summary DataFrame
    summary = pd.DataFrame(
        {
            "Metric": [
                "Total Observations",
                "Number of Events",
                "Number Censored",
                "Event Rate",
                "Median Survival Time",
                "Min Time",
                "Max Time",
                "Mean Time",
            ],
            "Value": [
                n_total,
                n_events,
                n_censored,
                f"{event_rate:.2%}",
                f"{median:.4f}",
                f"{time_min:.4f}",
                f"{time_max:.4f}",
                f"{time_mean:.4f}",
            ],
        }
    )

    return summary