Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 3 additions & 5 deletions src/stamp/modeling/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -1104,11 +1104,9 @@ def filter_complete_patient_data_(
)
}

_logger.info(
f"Total patients in clinical table: {len(patient_to_ground_truth)}\n"
f"Patients appearing in slide table: {len(patient_to_slides)}\n"
f"Final usable patients (complete data): {len(patients)}\n"
)
_logger.info("Total patients in clinical table: %d", len(patient_to_ground_truth))
_logger.info("Patients appearing in slide table: %d", len(patient_to_slides))
_logger.info("Final usable patients (complete data): %d", len(patients))
return patients


Expand Down
13 changes: 0 additions & 13 deletions src/stamp/modeling/models/cox.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@
# pylint: disable=C0103
# pylint: disable=C0301

import sys
import warnings

import torch
Expand Down Expand Up @@ -268,15 +267,3 @@ def neg_partial_log_likelihood(
)
)
return loss


if __name__ == "__main__":
import doctest

# Run doctest
results = doctest.testmod()
if results.failed == 0:
print("All tests passed.")
else:
print("Some doctests failed.")
sys.exit(1)
17 changes: 12 additions & 5 deletions src/stamp/statistics/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,8 +30,9 @@
)
from stamp.statistics.survival import _plot_km, _survival_stats_for_csv
from stamp.types import PandasLabel, Task
from stamp.utils.path import path_safe

__all__ = ["StatsConfig", "compute_stats_"]
__all__ = ["StatsConfig", "compute_stats_", "path_safe"]


__author__ = "Marko van Treeck, Minh Duc Nguyen"
Expand Down Expand Up @@ -146,7 +147,9 @@ def _compute_multitarget_classification_stats(
)

fig.tight_layout()
fig.savefig(output_dir / f"roc-curve_{target_label}={true_class}.svg")
fig.savefig(
output_dir / f"roc-curve_{path_safe(target_label)}={true_class}.svg"
)
plt.close(fig)

# Plot PRC curve
Expand All @@ -172,7 +175,9 @@ def _compute_multitarget_classification_stats(
)

fig.tight_layout()
fig.savefig(output_dir / f"pr-curve_{target_label}={true_class}.svg")
fig.savefig(
output_dir / f"pr-curve_{path_safe(target_label)}={true_class}.svg"
)
plt.close(fig)

# Compute aggregated statistics for all targets
Expand Down Expand Up @@ -292,7 +297,8 @@ def compute_stats_(
fig.tight_layout()
output_dir.mkdir(parents=True, exist_ok=True)
fig.savefig(
output_dir / f"roc-curve_{ground_truth_label}={true_class}.svg"
output_dir
/ f"roc-curve_{path_safe(ground_truth_label)}={true_class}.svg"
)
plt.close(fig)

Expand Down Expand Up @@ -321,7 +327,8 @@ def compute_stats_(

fig.tight_layout()
fig.savefig(
output_dir / f"pr-curve_{ground_truth_label}={true_class}.svg"
output_dir
/ f"pr-curve_{path_safe(ground_truth_label)}={true_class}.svg"
)
plt.close(fig)

Expand Down
18 changes: 14 additions & 4 deletions src/stamp/statistics/categorical.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
import scipy.stats as st
from sklearn import metrics

from stamp.utils.path import path_safe

__author__ = "Marko van Treeck"
__copyright__ = "Copyright (C) 2022-2025 Marko van Treeck"
__license__ = "MIT"
Expand Down Expand Up @@ -137,9 +139,13 @@ def categorical_aggregated_(
)

preds_df = pd.concat(preds_dfs).sort_index()
preds_df.to_csv(outpath / f"{ground_truth_label}_categorical-stats_individual.csv")
preds_df.to_csv(
outpath / f"{path_safe(ground_truth_label)}_categorical-stats_individual.csv"
)
stats_df = _aggregate_categorical_stats(preds_df.reset_index())
stats_df.to_csv(outpath / f"{ground_truth_label}_categorical-stats_aggregated.csv")
stats_df.to_csv(
outpath / f"{path_safe(ground_truth_label)}_categorical-stats_aggregated.csv"
)


def categorical_aggregated_multitarget_(
Expand Down Expand Up @@ -181,11 +187,15 @@ def categorical_aggregated_multitarget_(

# Concatenate and save individual stats for this target
preds_df = pd.concat(preds_dfs).sort_index()
preds_df.to_csv(outpath / f"{target_label}_categorical-stats_individual.csv")
preds_df.to_csv(
outpath / f"{path_safe(target_label)}_categorical-stats_individual.csv"
)

# Aggregate stats for this target
stats_df = _aggregate_categorical_stats(preds_df.reset_index())
stats_df.to_csv(outpath / f"{target_label}_categorical-stats_aggregated.csv")
stats_df.to_csv(
outpath / f"{path_safe(target_label)}_categorical-stats_aggregated.csv"
)

# Store for summary
all_target_stats[target_label] = stats_df
Expand Down
10 changes: 8 additions & 2 deletions src/stamp/statistics/regression.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@
import scipy.stats as st
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score

from stamp.utils.path import path_safe


def _regression(preds_df: pd.DataFrame, target_label: str) -> pd.Series:
"""Compute regression metrics for one prediction table."""
Expand Down Expand Up @@ -107,10 +109,14 @@ def regression_aggregated_(

# Save individual stats and aggregate
stats_df = pd.DataFrame(stats).transpose()
stats_df.to_csv(outpath / f"{ground_truth_label}_regression-stats_individual.csv")
stats_df.to_csv(
outpath / f"{path_safe(ground_truth_label)}_regression-stats_individual.csv"
)

mean = stats_df.mean(numeric_only=True)
sem = stats_df.sem(numeric_only=True)
lower, upper = st.t.interval(0.95, len(stats_df) - 1, loc=mean, scale=sem)
agg = pd.DataFrame({"mean": mean, "95%_low": lower, "95%_high": upper})
agg.to_csv(outpath / f"{ground_truth_label}_regression-stats_aggregated.csv")
agg.to_csv(
outpath / f"{path_safe(ground_truth_label)}_regression-stats_aggregated.csv"
)
6 changes: 6 additions & 0 deletions src/stamp/utils/path.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
"""Small path-related helpers shared across the project."""


def path_safe(label: str) -> str:
"""Replace path separators so labels are safe as filename parts."""
return label.replace("/", "_").replace("\\", "_")
Loading