viral-images / features /saliency.py
Babajaan's picture
Full Viral Images v1.0 implementation - all modules and configs
6ceaa94 verified
Raw
History Blame Contribute Delete
2.2 kB
"""Saliency map computation using pretrained MSI-Net (fallback to heuristic)."""
import numpy as np
from PIL import Image
from scipy.stats import entropy as scipy_entropy
from typing import Dict, Optional
def _compute_heuristic_saliency(img: Image.Image) -> np.ndarray:
"""
Fallback saliency using luminance contrast and color variance.
Returns a pseudo-saliency map in [0,1].
"""
arr = np.array(img.convert("RGB")).astype(np.float32)
gray = arr.mean(axis=2)
# Compute local contrast via gradient magnitude
dx = np.abs(np.diff(gray, axis=1, append=gray[:, -1:]))
dy = np.abs(np.diff(gray, axis=0, append=gray[-1:, :]))
grad = np.sqrt(dx**2 + dy**2)
# Normalize
sal = grad / (grad.max() + 1e-8)
return sal.astype(np.float32)
def compute_saliency_features(img: Image.Image, saliency_map: Optional[np.ndarray] = None) -> Dict[str, float]:
"""
Compute saliency-based features from an image.
If MSI-Net fails or is not loaded, falls back to heuristic.
"""
if saliency_map is None:
saliency_map = _compute_heuristic_saliency(img)
sal_flat = saliency_map.ravel()
sal_probs = sal_flat / (sal_flat.sum() + 1e-8)
features = {}
# 1. Peak saliency
features["peak_saliency"] = float(saliency_map.max())
# 2. Mean saliency
features["mean_saliency"] = float(saliency_map.mean())
# 3. Saliency entropy (spatial spread)
max_entropy = np.log(len(sal_flat) + 1e-8)
features["saliency_entropy"] = float(scipy_entropy(sal_probs + 1e-10)) / max_entropy
# 4. Top-20% fraction
threshold = 0.8 * features["peak_saliency"]
features["top20_fraction"] = float((saliency_map > threshold).sum()) / (saliency_map.size + 1e-8)
# 5. Center saliency (central third)
H, W = saliency_map.shape
h3, w3 = max(1, H // 3), max(1, W // 3)
center_zone = saliency_map[h3:2*h3, w3:2*w3]
features["center_saliency"] = float(center_zone.mean())
# 6. Saliency uniformity (lower = more varied = more interesting)
features["saliency_uniformity"] = float(saliency_map.std() / (features["mean_saliency"] + 1e-8))
return features