From a57fe0cba3b96bee59470f3299868ad8bc9e9bdf Mon Sep 17 00:00:00 2001 From: "Salvador E. Tropea" Date: Fri, 7 Nov 2025 13:45:29 -0300 Subject: [PATCH] [SaliencyEvaluationMetrics] Better computations - All moved outside the main source - Documented - Most replaced by proven versions --- docs/saliency_metrics.md | 1 - src/nodes/e_measure.py | 123 ++++++++++++++++++++++++ src/nodes/f_measure.py | 179 +++++++++++++++++++++++++++++++++++ src/nodes/nodes_img.py | 150 ++++-------------------------- src/nodes/s_measure.py | 196 +++++++++++++++++++++++++++++++++++++++ 5 files changed, 514 insertions(+), 135 deletions(-) create mode 100644 src/nodes/e_measure.py create mode 100644 src/nodes/f_measure.py create mode 100644 src/nodes/s_measure.py diff --git a/docs/saliency_metrics.md b/docs/saliency_metrics.md index 17e00a2..51b89f0 100644 --- a/docs/saliency_metrics.md +++ b/docs/saliency_metrics.md @@ -27,7 +27,6 @@ This document details the quantitative metrics used to evaluate the performance * **What it Measures**: The F-measure is the harmonic mean of Precision and Recall, providing a score that balances the two. In saliency evaluation, the continuous prediction map is converted to a binary map using a series of thresholds (from 0 to 255). The F-measure is calculated for each threshold, and the **maximum** value obtained across all thresholds is reported. This adaptive thresholding makes the metric robust to models that produce well-shaped but poorly-calibrated (e.g., generally too dark or bright) saliency maps. * **Interpretation**: - * **Range**:. * **Higher is better**. A score of **1** represents a perfect balance of precision and recall at the optimal threshold. * **Relevance and Justification**: Unlike the pixel-level MAE, the F-measure is region-based. It evaluates how well the *shape* of the predicted salient region aligns with the ground truth. By finding the optimal threshold for a given prediction, it fairly assesses the quality of the saliency map's structure, forgiving issues with overall intensity. The standard beta-squared value ($β^2$) is set to **0.3** to weigh precision more heavily than recall, as proposed by the authors of the foundational paper below. diff --git a/src/nodes/e_measure.py b/src/nodes/e_measure.py new file mode 100644 index 0000000..9ffd419 --- /dev/null +++ b/src/nodes/e_measure.py @@ -0,0 +1,123 @@ +import torch +from typing import Tuple +# Use a safe epsilon for numerical stability +EPS = 1e-8 + + +def get_e_measure( + pred: torch.Tensor, + gt: torch.Tensor, + num_thresholds: int = 255, + chunk_size: int = 16 +) -> Tuple[float, float, float, torch.Tensor]: + """ + Calculates the E-measure scores using a memory-efficient chunking strategy. + + This implementation is fully vectorized within chunks to maintain high performance + while ensuring that memory usage remains low and predictable, making it suitable + for high-resolution images. + + Args: + pred (torch.Tensor): The continuous prediction mask (normalized to [0, 1]). + gt (torch.Tensor): The binary ground truth mask (values are 0 or 1). + num_thresholds (int): The number of thresholds to evaluate. Defaults to 255. + chunk_size (int): The number of thresholds to process in a single batch. + Lower this value if you encounter VRAM issues. Defaults to 16. + + Returns: + A tuple containing: + - float: The mean E-measure score across all thresholds. + - float: The maximum E-measure score across all thresholds. + - float: The adaptive E-measure score. + - torch.Tensor: A 1D tensor with the E-measure score for each threshold. + """ + + # 1. --- Calculate scores for all thresholds using chunking --- + + # Create a 1D tensor of thresholds from 0 to almost 1 + thlist = torch.linspace(0, 1 - 1e-10, num_thresholds, device=pred.device) + all_scores = [] + + # Process thresholds in memory-efficient chunks + for i in range(0, num_thresholds, chunk_size): + # Get the current chunk of thresholds + th_chunk = thlist[i:i + chunk_size] + + # Vectorized operation on the smaller chunk. + # Reshape pred to (1, H, W) and thlist to (N, 1, 1). + # Broadcasting (>=) creates a binarized prediction for each threshold. + # The result `binarized_preds` has a shape of (chunk_size, H, W). + binarized_preds_chunk = (pred.unsqueeze(0) >= th_chunk.view(-1, 1, 1)).to(pred.dtype) + + # Calculate scores for the current chunk + scores_chunk = e_calculate_enhanced_scores(binarized_preds_chunk, gt) + all_scores.append(scores_chunk) + + # Combine the scores from all chunks + scores = torch.cat(all_scores) + + # 2. --- Calculate the single adaptive score --- + + # Calculate the adaptive threshold, clamping at 1.0 to be safe. + adaptive_th = torch.clamp(2 * pred.mean(), max=1.0) + # Binarize the prediction with this single threshold + adaptive_pred_binarized = (pred >= adaptive_th).to(pred.dtype) + # Reuse the same helper function by adding a temporary batch dimension + adaptive_score_tensor = e_calculate_enhanced_scores(adaptive_pred_binarized.unsqueeze(0), gt) + + # 3. --- Return the final results --- + + return ( + scores.mean().item(), + scores.max().item(), + adaptive_score_tensor.item(), + scores + ) + + +def e_calculate_enhanced_scores(binarized_preds: torch.Tensor, gt: torch.Tensor) -> torch.Tensor: + """ + Helper function to compute E-measure scores for a batch of binarized predictions. + This function is fully vectorized. + + Args: + binarized_preds (torch.Tensor): A tensor of binarized predictions, with shape + (N, H, W), where N is the number of thresholds. + gt (torch.Tensor): The single binary ground truth mask, with shape (H, W). + + Returns: + torch.Tensor: A 1D tensor of shape (N,) containing the E-measure score + for each binarized prediction. + """ + # Handle the edge case where the ground truth is all black + if torch.mean(gt) == 0.0: + # The score is based on how much of the prediction is also black. + # Original formula: sum(1 - y_pred_th) / (numel - 1) + # This is equivalent to numel * (1 - mean) / (numel - 1) + enhanced = 1 - binarized_preds + + # Handle the edge case where the ground truth is all white + elif torch.mean(gt) == 1.0: + # The score is based on how much of the prediction is also white. + # Original formula: sum(y_pred_th) / (numel - 1) + enhanced = binarized_preds + + # Normal case with a mixed ground truth + else: + # Demean the ground truth. `gt_demeaned` has shape (H, W). + gt_demeaned = gt - gt.mean() + + # Demean the binarized predictions. + # `mean` is calculated over spatial dims (H, W), keeping the threshold dim. + # `fm` (foreground map) will have shape (N, H, W). + fm = binarized_preds - binarized_preds.mean(dim=[-2, -1], keepdim=True) + + # The demeaned GT will be broadcasted to match the shape of `fm`. + align_matrix = 2 * gt_demeaned * fm / (gt_demeaned.square() + fm.square() + EPS) + enhanced = (align_matrix + 1).square() / 4 + + # Calculate the final score for each threshold by summing over the spatial dimensions. + # The denominator (y.numel() - 1) is a quirk from the original paper's code. + scores = torch.sum(enhanced, dim=[-2, -1]) / (gt.numel() - 1 + EPS) + + return scores diff --git a/src/nodes/f_measure.py b/src/nodes/f_measure.py new file mode 100644 index 0000000..9ca6557 --- /dev/null +++ b/src/nodes/f_measure.py @@ -0,0 +1,179 @@ +import torch +import numpy as np +import scipy +# Epsilon: small value to avoid "divide by 0" errors +EPS = 1e-8 + + +def get_f_measure(pred: torch.Tensor, gt: torch.Tensor, beta2: float = 0.3) -> float: + """ + Calculates the maximum F-measure score for a continuous prediction against a binary ground truth. + + The F-measure evaluates the balance between precision and recall. Since the prediction + is a continuous map (0.0 to 1.0), this function iterates through 256 possible + thresholds to binarize the prediction. It calculates the F-measure for each + threshold and returns the highest (max) score found. This provides a fair + evaluation of the prediction's structural quality, independent of its overall brightness. + + Args: + pred (torch.Tensor): The continuous prediction mask (normalized to [0, 1]). + gt (torch.Tensor): The binary ground truth mask (values are 0 or 1). + beta2 (float): The beta-squared value for the F-measure. The standard value of 0.3 + is used to weigh precision more heavily than recall. Defaults to 0.3. + + Returns: + float: The maximum F-measure score found across all thresholds. + """ + # Initialize f_max to store the highest F-measure score found so far. + f_max = 0.0 + + # Iterate through 256 evenly spaced thresholds from 0.0 to 1.0. + # This corresponds to testing every possible 8-bit grayscale value as the cutoff. + # The thresholds tensor is created on the same device as the input for efficiency. + for threshold in torch.linspace(0, 1, 256, device=gt.device): + # Binarize the continuous prediction map using the current threshold. + # Pixels >= threshold become 1.0 (positive), and others become 0.0 (negative). + pred_binary = (pred >= threshold).float() + + # Calculate True Positives (TP): pixels that are positive in both the prediction and ground truth. + # Element-wise multiplication results in 1 only where both are 1. + tp = (pred_binary * gt).sum() + + # Optimization: If there are no true positives, the F-measure will be 0. + # We can skip the rest of the calculations for this threshold. + if tp == 0: + continue + + # Calculate Precision = TP / (TP + FP). + # The sum of `pred_binary` gives the total number of predicted positives (TP + FP). + precision = tp / (pred_binary.sum() + EPS) + + # Calculate Recall = TP / (TP + FN). + # The sum of `gt` gives the total number of actual positives (TP + FN). + recall = tp / (gt.sum() + EPS) + + # Calculate the F-beta score using the computed precision and recall. + # The beta^2=0.3 value is standard in saliency detection literature. + f_beta = (1 + beta2) * precision * recall / (beta2 * precision + recall + EPS) + + # Update f_max if the F-beta score for the current threshold is the highest yet. + # .item() extracts the single float value from the 0-dimensional tensor. + if f_beta > f_max: + f_max = f_beta.item() + + # After checking all thresholds, return the maximum score found. + return f_max + + +def get_weighted_f_measure(pred: torch.Tensor, gt: torch.Tensor, beta2: float = 0.3) -> float: + """ + Calculates the true research-grade Weighted F-measure (F_beta^w). + + This is a faithful and optimized port of the reference algorithm from the + "How to Evaluate Foreground Maps?" paper, ensuring verifiable results. + The implementation performs pre-checks on the GPU for efficiency before + transferring data to the CPU for SciPy-based calculations. + + Args: + pred (torch.Tensor): The continuous prediction mask (normalized to [0, 1]). + gt (torch.Tensor): The binary ground truth mask (values are 0 or 1). + beta2 (float): The beta-squared value for the F-measure. Defaults to 0.3. + + Returns: + float: The final Weighted F-measure score. + """ + # --- Step 1: GPU-side Pre-checks for Efficiency --- + + # Handle the edge case of an all-black ground truth on the GPU. + # If the ground truth is empty, the score is 1 minus the mean of the prediction. + # A perfect prediction (all black) would yield a score of 1. + # This avoids the expensive CPU transfer and SciPy calculations entirely. + if torch.mean(gt) == 0.0: + return (1.0 - pred.mean()).item() + + # --- Step 2: Data Transfer to CPU for NumPy/SciPy Processing --- + + # Move tensors to the CPU and convert to NumPy arrays. SciPy functions require this. + # Squeezing removes any singleton channel dimensions (e.g., from [1, H, W] to [H, W]). + gt_np = gt.squeeze().cpu().numpy() + pred_np = pred.squeeze().cpu().numpy() + + # --- Step 3: Creation of Core Error and Dependency Maps --- + + # Create boolean masks for foreground (gt_mask) and background (not_gt_mask) regions. + # np.isclose is used for safe floating-point comparison. + gt_mask = np.isclose(gt_np, 1) + not_gt_mask = np.logical_not(gt_mask) + + # Calculate the initial, simple absolute pixel-wise error map. + E = np.abs(pred_np - gt_np) + + # Calculate the Euclidean Distance Transform on the INVERTED mask. + # For each background pixel, `dist` will be its distance to the nearest foreground pixel. + # `idx` will store the coordinates of that nearest foreground pixel. This is key for the next step. + # Note: The original implementation uses scipy.ndimage.morphology.distance_transform_edt + # which is an alias for scipy.ndimage.distance_transform_edt. + dist, idx = scipy.ndimage.morphology.distance_transform_edt(not_gt_mask, return_indices=True) + + # --- Step 4: Pixel Dependency Map (Et -> EA -> min_E_EA) --- + + # Create the "Pixel Dependency" map, starting with a copy of the original error. + Et = np.array(E) + + # This is the crucial step for pixel dependency. For every background pixel, + # its error value is replaced with the error value of its NEAREST foreground pixel. + # This propagates the error from the foreground edges into the background. + Et[not_gt_mask] = E[idx[0, not_gt_mask], idx[1, not_gt_mask]] + + # Smooth the dependency-aware error map with a Gaussian filter. + # This creates a soft, blurred error field around the object. + sigma = 5.0 + EA = scipy.ndimage.gaussian_filter(Et, sigma=sigma, truncate=3 / sigma, mode='constant', cval=0.0) + + # The final error for any FOREGROUND pixel is the MINIMUM of its original error (E) + # and the new smoothed, dependency-aware error (EA). + # This prevents unfairly penalizing small errors right at the boundary of a correct prediction. + # The `where=gt_mask` argument ensures this operation only applies to the foreground. + min_E_EA = np.minimum(E, EA, where=gt_mask, out=np.array(E)) + + # --- Step 5: Pixel Importance Map (B) --- + + # Create the "Pixel Importance" map, which starts as a uniform map of ones. + B = np.ones(gt_np.shape) + + # For each BACKGROUND pixel, assign an importance weight based on its distance + # from the foreground. The weight is calculated using a non-linear function + # that decreases as the distance increases. This makes errors near the object more important. + B[not_gt_mask] = 2 - np.exp(np.log(1 - 0.5) / 5 * dist[not_gt_mask]) + + # The final Weighted Error Map is the element-wise product of the dependency-aware + # error map and the pixel importance map. + Ew = min_E_EA * B + + # --- Step 6: Final Metric Computation --- + + # Get a small machine epsilon value for numerically stable division. + eps = np.spacing(1) + + # Calculate Weighted True Positives (TPw) and False Positives (FPw) from the Ew map. + # TPw = Sum of foreground weights - Sum of weighted errors in the foreground. + # FPw = Sum of weighted errors in the background. + TPw = np.sum(gt_np) - np.sum(Ew[gt_mask]) + FPw = np.sum(Ew[not_gt_mask]) + + # Calculate Weighted Recall (R) and Weighted Precision (P). + # These definitions are specific to this metric's formulation. + R = 1 - np.mean(Ew[gt_mask]) # Weighed Recall + P = TPw / (eps + TPw + FPw) # Weighted Precision + + # The final Weighted F-measure (Q) is calculated using the standard formula + # with the newly computed weighted Precision and Recall. + # beta2 = 0.3 is standard, weighing precision more heavily than recall. + # Q = 2 * (R * P) / (eps + R + P) # Beta=1 + Q = (1 + beta2) * (R * P) / (eps + R + (beta2 * P)) + + # Raise an error if the result is Not a Number (NaN), indicating a potential issue. + if np.isnan(Q): + raise ValueError("Weighted F-measure resulted in NaN") + + return Q diff --git a/src/nodes/nodes_img.py b/src/nodes/nodes_img.py index 2fab813..20be0fd 100644 --- a/src/nodes/nodes_img.py +++ b/src/nodes/nodes_img.py @@ -29,6 +29,9 @@ from typing import Optional # We are the main source, so we use the main_logger from . import main_logger from .helpers import load_image_wrapper, load_images_wrapper, save_image, upscale, upscale_comfy +from .s_measure import get_s_measure +from .e_measure import get_e_measure +from .f_measure import get_f_measure, get_weighted_f_measure try: from folder_paths import get_input_directory, get_output_directory except ModuleNotFoundError: @@ -659,109 +662,6 @@ class MaskDifference: return (diff_image_bhwc,) -# --- Helper functions for advanced metrics --- -# These implementations are PyTorch adaptations of common saliency evaluation libraries. -# Credit to the original authors of S-measure, E-measure, and Weighted F-measure. - -def _get_s_measure(pred, gt): - alpha = 0.5 - y = gt.mean() - if y == 0: - x = pred.mean() - q = 1.0 - x - elif y == 1: - x = pred.mean() - q = x - else: - # gt is assumed to be binary - q = alpha * _object(pred, gt) + (1 - alpha) * _region(pred, gt) - if q < 0: - q = torch.tensor([0.0], device=pred.device) - return q - - -def _object(pred, gt): - fg = torch.where(gt == 0, torch.zeros_like(pred), pred) - bg = torch.where(gt == 1, torch.zeros_like(pred), 1 - pred) - o_fg = _object_calc(fg, gt) - o_bg = _object_calc(bg, 1 - gt) - u = gt.mean() - q = u * o_fg + (1 - u) * o_bg - return q - - -def _object_calc(pred, gt): - x = pred.mean() - # sigma_x = pred.std() - score = 2.0 * x / (x**2 + 1.0 + 1e-8) - return score - - -def _region(pred, gt): - [y, x] = torch.where(gt == 1) - if len(y) == 0 or len(x) == 0: - return torch.tensor(0.0) - - y_bar, x_bar = y.float().mean(), x.float().mean() - - gt_w = torch.where(gt == 0, gt.float(), 1-gt.float()) - gt_w[y_bar.long(), x_bar.long()] = 1 - - gt_w_sum = gt_w.sum() - if gt_w_sum > 0: - gt_w = gt_w / gt_w_sum - else: # Handle case where sum is zero - gt_w.fill_(1.0 / (gt.shape[0] * gt.shape[1])) - - parts = gt_w.unique() - - if len(parts) == 1: - return (1-pred).mean() if parts[0] == 0 else pred.mean() - - w_pred = pred * gt_w - return w_pred.sum() - - -def _get_e_measure(pred, gt): - # gt is assumed to be binary - pred = (pred - pred.mean()) / (pred.std() + 1e-8) - gt = (gt - gt.mean()) / (gt.std() + 1e-8) - - align_matrix = 2 * gt * pred / (gt * gt + pred * pred + 1e-8) - enhanced = (align_matrix + 1)**2 / 4 - - score = torch.mean(enhanced) - return score - - -def _get_weighted_f_measure(pred, gt): - # gt is assumed to be binary - - # Implementation based on https://github.com/wenguanwang/SODsurvey/ - # Generates a weight map that gives more importance to pixels near the center. - center_w = torch.ones_like(gt) - h, w = gt.shape - y, x = torch.meshgrid(torch.arange(h, device=gt.device), torch.arange(w, device=gt.device), indexing="ij") - - center_w = 1 - 0.5 * (torch.abs(y - (h-1)/2) / ((h-1)/2) + torch.abs(x - (w-1)/2) / ((w-1)/2)) - - tp = center_w * (pred * gt) - fp = center_w * (pred * (1-gt)) - fn = center_w * ((1-pred) * gt) - - # Add a small epsilon to avoid division by zero - eps = 1e-6 - - prec = tp.sum() / (tp.sum() + fp.sum() + eps) - recall = tp.sum() / (tp.sum() + fn.sum() + eps) - - # Using beta^2 = 0.3 as is standard. - beta2 = 0.3 - f_beta = (1 + beta2) * prec * recall / (beta2 * prec + recall + eps) - - return f_beta - - class SaliencyEvaluationMetrics: @classmethod def INPUT_TYPES(s): @@ -817,7 +717,6 @@ class SaliencyEvaluationMetrics: # --- Initialize accumulators for metrics --- mae_total, f_measure_max_total, s_measure_total, e_measure_total, weighted_f_total = 0, 0, 0, 0, 0 - eps = 1e-6 all = [] for i in range(batch_size): @@ -827,7 +726,7 @@ class SaliencyEvaluationMetrics: # 1. Mean Absolute Error (MAE) if mae_enable: - mae = torch.mean(torch.abs(pred_i - gt_i)) + mae = torch.mean(torch.abs(pred_i - gt_i)).item() logger.debug(f"MAE: {mae}") mae_total += mae res['mae'] = mae @@ -838,45 +737,28 @@ class SaliencyEvaluationMetrics: # 2. Max F-measure if max_f_mes_enable: - f_max = 0.0 - for threshold in torch.linspace(0, 1, 256, device=device): - pred_binary = (pred_i >= threshold).float() - - tp = (pred_binary * gt_binary).sum() - - if tp == 0: - continue - - precision = tp / (pred_binary.sum() + eps) - recall = tp / (gt_binary.sum() + eps) - - # Using beta^2 = 0.3 as is standard. - beta2 = 0.3 - f_beta = (1 + beta2) * precision * recall / (beta2 * precision + recall + eps) - - if f_beta > f_max: - f_max = f_beta + f_max = get_f_measure(pred_i, gt_binary) f_measure_max_total += f_max logger.debug(f"F_max: {f_max}") res['max_f_mes'] = f_max # 3. S-measure if s_mes_enable: - s_measure = _get_s_measure(pred_i, gt_binary) + s_measure = get_s_measure(pred_i, gt_binary) s_measure_total += s_measure logger.debug(f"S: {s_measure}") res['s_mes'] = s_measure # 4. E-measure if e_mes_enable: - e_measure = _get_e_measure(pred_i, gt_binary) - e_measure_total += e_measure - logger.debug(f"E: {e_measure}") - res['e_mes'] = e_measure + e_mean, e_max, e_adp, _ = get_e_measure(pred_i, gt_binary) + e_measure_total += e_mean + logger.debug(f"E: {e_mean} {e_max} {e_adp}") + res['e_mes'] = e_mean # 5. Weighted F-measure if wf_mes_enable: - wf = _get_weighted_f_measure(pred_i, gt_binary) + wf = get_weighted_f_measure(pred_i, gt_binary) weighted_f_total += wf logger.debug(f"wF: {wf}") res['wf_mes'] = wf @@ -884,11 +766,11 @@ class SaliencyEvaluationMetrics: all.append(res) # --- Average metrics over the batch --- - mae_avg = mae_total.item() / batch_size - f_measure_avg = f_measure_max_total.item() / batch_size - s_measure_avg = s_measure_total.item() / batch_size - e_measure_avg = e_measure_total.item() / batch_size - weighted_f_avg = weighted_f_total.item() / batch_size + mae_avg = mae_total / batch_size + f_measure_avg = f_measure_max_total / batch_size + s_measure_avg = s_measure_total / batch_size + e_measure_avg = e_measure_total / batch_size + weighted_f_avg = weighted_f_total / batch_size return (all, mae_avg, f_measure_avg, s_measure_avg, e_measure_avg, weighted_f_avg) diff --git a/src/nodes/s_measure.py b/src/nodes/s_measure.py new file mode 100644 index 0000000..cea669e --- /dev/null +++ b/src/nodes/s_measure.py @@ -0,0 +1,196 @@ +import torch +# Epsilon: small value to avoid "divide by 0" errors +EPS = 1e-8 + + +def get_s_measure(pred: torch.Tensor, gt: torch.Tensor, alpha: float = 0.5) -> float: + """ + Calculates the S-measure for a given prediction and ground truth mask. + The S-measure evaluates structural similarity, combining object-aware and + region-aware metrics. + + Args: + pred (torch.Tensor): The continuous prediction mask (normalized to [0, 1]). + gt (torch.Tensor): The binary ground truth mask (values are 0 or 1). + alpha (float): The weight for balancing object-aware vs. region-aware scores. + Defaults to 0.5. + + Returns: + float: The final S-measure score. + """ + # gt is assumed to be binary + y = gt.mean() + if y == 0: + # If the ground truth is all black, the score is the inverse of the prediction's average. + # A perfect prediction would also be all black (mean=0), yielding a score of 1. + x = pred.mean() + q = 1.0 - x + elif y == 1: + # If the ground truth is all white, the score is simply the prediction's average. + # A perfect prediction would be all white (mean=1), yielding a score of 1. + x = pred.mean() + q = x + else: + # For a mixed ground truth, balance the object and region scores. + q = alpha * s_object(pred, gt) + (1 - alpha) * s_region(pred, gt) + # Ensure the score is non-negative. + if q < 0: + return 0 + + return q.item() + + +def s_object(pred: torch.Tensor, gt: torch.Tensor) -> torch.Tensor: + """ + Calculates the object-aware structural similarity score. + This measures the similarity for foreground and background regions separately + and combines them based on the foreground's size. + """ + # Create a map of foreground-only predictions + fg = torch.where(gt == 0, torch.zeros_like(pred), pred) + # Create a map of background-only predictions (inverted) + bg = torch.where(gt == 1, torch.zeros_like(pred), 1 - pred) + + # Calculate scores for each region + o_fg = s_object_calc(fg, gt) + o_bg = s_object_calc(bg, 1 - gt) + + # Combine scores based on foreground area + u = gt.mean() # The ratio of foreground pixels + q = u * o_fg + (1 - u) * o_bg + return q + + +def s_object_calc(pred: torch.Tensor, gt: torch.Tensor) -> torch.Tensor: + """ + Helper function to compute the similarity score for a given region (FG or BG). + """ + # Select the prediction pixels corresponding to the region of interest in the ground truth + region_pixels = pred[gt == 1] + + # If the region is empty, the score is undefined, but 0 is a safe return. + if region_pixels.numel() == 0: + return torch.tensor(0.0, device=pred.device) + + x = region_pixels.mean() + + # Note: The original paper's formula uses variance (sigma**2), but many public + # implementations use std dev (sigma). We follow this implementation's logic. + sigma_x = region_pixels.std() + + # A score that rewards high mean (x) but penalizes high variance (sigma_x). + # It is maximized when x=1 and sigma_x=0. + score = 2.0 * x / (x**2 + 1.0 + sigma_x + EPS) + return score + + +def s_region(pred: torch.Tensor, gt: torch.Tensor) -> torch.Tensor: + """ + Calculates the region-aware structural similarity score. + This divides the image into four quadrants based on the ground truth's + centroid and computes a weighted SSIM score. + """ + # Find the centroid (center of mass) of the ground truth mask + X, Y = s_centroid(gt) + + # Divide the ground truth and prediction into four quadrants based on the centroid + gt_parts = s_divide_tensor(gt, X, Y) + pred_parts = s_divide_tensor(pred, X, Y) + + # Calculate the area weights for each quadrant + w = s_calculate_weights(gt.shape, X, Y) + + # Calculate SSIM for each quadrant and combine the quadrant scores using their area weights + return sum([w[i] * s_ssim(pred_parts[i], gt_parts[i]) for i in range(4)]) + + +def s_centroid(gt: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """ + Calculates the centroid (center of mass) of a binary mask. + """ + rows, cols = gt.shape[-2:] + total = gt.sum() + + # Handle the edge case of an all-black mask + if total == 0: + X = torch.tensor(round(cols / 2), device=gt.device) + Y = torch.tensor(round(rows / 2), device=gt.device) + else: + # Create coordinate ranges directly on the target device, avoiding CPU-GPU transfer. + i_coords = torch.arange(cols, device=gt.device, dtype=gt.dtype) + j_coords = torch.arange(rows, device=gt.device, dtype=gt.dtype) + + # Calculate weighted average of coordinates + X = torch.round((gt.sum(dim=0) * i_coords).sum() / total) + Y = torch.round((gt.sum(dim=1) * j_coords).sum() / total) + + return X.long(), Y.long() + + +def s_divide_tensor(tensor: torch.Tensor, X: torch.Tensor, Y: torch.Tensor) -> tuple: + """ + Divides a tensor into four quadrants based on a pivot point (X, Y). + This is memory-efficient as it returns views, not copies. + """ + h, w = tensor.shape[-2:] + LT = tensor[..., :Y, :X] + RT = tensor[..., :Y, X:] + LB = tensor[..., Y:, :X] + RB = tensor[..., Y:, X:] + return LT, RT, LB, RB + + +def s_calculate_weights(shape: tuple, X: torch.Tensor, Y: torch.Tensor) -> tuple: + """Calculates the proportional area of the four quadrants.""" + h, w = shape[-2:] + area = h * w + + # Ensure coordinates are float for division + Xf = X.float() + Yf = Y.float() + + w1 = Xf * Yf / area + w2 = (w - Xf) * Yf / area + w3 = Xf * (h - Yf) / area + w4 = 1.0 - w1 - w2 - w3 # More stable calculation for the last weight + + return (w1, w2, w3, w4) + + +def s_ssim(pred: torch.Tensor, gt: torch.Tensor) -> float: + """ + Computes a custom structural similarity (SSIM-like) score between two tensors. + """ + # If a quadrant is empty, its contribution to similarity is ambiguous. + # Returning 0 is a safe choice, but 1 could also be argued if gt is also empty. + if pred.numel() == 0 or gt.numel() == 0: + return 0.0 + + h, w = pred.shape[-2:] + N = h * w + + # Means + x = pred.mean() + y = gt.mean() + + # Variances and Covariance (using unbiased estimator N-1) + # This is numerically safer than calculating std dev separately. + sigma_x2 = ((pred - x) * (pred - x)).sum() / (N - 1 + EPS) + sigma_y2 = ((gt - y) * (gt - y)).sum() / (N - 1 + EPS) + sigma_xy = ((pred - x) * (gt - y)).sum() / (N - 1 + EPS) + + # Numerator and denominator of the SSIM formula + alpha = 4 * x * y * sigma_xy + beta = (x*x + y*y) * (sigma_x2 + sigma_y2) + + # Handle special cases for stability, as defined in the original code + if alpha != 0: + Q = alpha / (beta + EPS) + elif alpha == 0 and beta == 0: + # If both inputs are flat and identical, similarity is perfect. + Q = 1.0 + else: + # If numerator is 0 but denominator isn't, similarity is 0. + Q = 0 + + return Q