From fecff541e6353e94cc121b05d74a75b50de96060 Mon Sep 17 00:00:00 2001 From: "Salvador E. Tropea" Date: Sat, 29 Nov 2025 19:02:46 -0300 Subject: [PATCH] [Added] Automatic batch computation for E-measure It went OoM for an 8800x7767 image --- src/nodes/e_measure.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/nodes/e_measure.py b/src/nodes/e_measure.py index 5686c42..bbcf72c 100644 --- a/src/nodes/e_measure.py +++ b/src/nodes/e_measure.py @@ -7,7 +7,7 @@ def get_e_measure( pred: torch.Tensor, gt: torch.Tensor, num_thresholds: int = F_POINTS, - chunk_size: int = 16 + chunk_size: int = -1 ) -> Tuple[float, float, float, torch.Tensor, torch.Tensor]: """ Calculates the E-measure scores using a memory-efficient chunking strategy. @@ -38,6 +38,10 @@ def get_e_measure( thlist = torch.linspace(0, 1 - 1e-10, num_thresholds, device=pred.device) all_scores = [] + if chunk_size == -1: + chunk_size = round(64 / (pred.shape[0] * pred.shape[1] / (1<<20))) + chunk_size = min(max(1, chunk_size), num_thresholds) + # Process thresholds in memory-efficient chunks for i in range(0, num_thresholds, chunk_size): # Get the current chunk of thresholds