[SaliencyEvaluationMetrics] Better computations

- All moved outside the main source
- Documented
- Most replaced by proven versions
This commit is contained in:
Salvador E. Tropea
2025-11-07 13:45:29 -03:00
parent a2ca2fe876
commit a57fe0cba3
5 changed files with 514 additions and 135 deletions
-1
View File
@@ -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.
+123
View File
@@ -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
+179
View File
@@ -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
+16 -134
View File
@@ -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)
+196
View File
@@ -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