[SaliencyEvaluationMetrics] Better computations
- All moved outside the main source - Documented - Most replaced by proven versions
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user