Whyx-PROmpTea / src /background_analyzer.py
ArtShumov's picture
feat: add anatomy, expression, and background analyzers
536a29a
Raw
History Blame Contribute Delete
5.26 kB
"""Background analysis: location, atmosphere, and lighting from image + WD14 tags.
Uses WD14 tags as the primary source (they already detect many background
and lighting tags) and supplements with lightweight color-based heuristics
that run on the image border regions.
All outputs are Danbooru-format tags (underscores).
"""
from __future__ import annotations
import numpy as np
# Location-related tags from tag pools
_BG_LOCATION_KEYWORDS = {
"indoor": {"indoor", "room", "bedroom", "kitchen", "classroom", "library",
"indoors", "interior", "hallway", "bathroom", "toilet", "office",
"shop", "store", "staircase", "elevator", "laboratory"},
"outdoor": {"outdoor", "sky", "cloud", "sun", "sunset", "sunrise",
"nature", "forest", "tree", "mountain", "beach", "sea",
"ocean", "river", "lake", "waterfall", "field", "meadow",
"desert", "ruins", "castle", "village", "street", "city",
"rooftop", "bridge", "park", "garden", "path", "road"},
"special": {"underwater", "space", "cyberpunk", "fantasy", "surreal"},
}
_BG_ATMOSPHERE_KEYWORDS = {
"weather": {"rain", "snow", "fog", "mist", "wind", "storm", "cloudy",
"clear", "starry", "aurora"},
"mood": {"atmospheric", "ethereal", "dreamlike", "romantic",
"melancholic", "serene", "dark", "moody", "dramatic",
"cozy", "calm", "chaotic", "peaceful"},
"effects": {"petals", "falling leaves", "fireflies", "sparkles",
"glow", "light rays", "particles"},
}
_BG_LIGHTING_KEYWORDS = {
"natural": {"sunlight", "sunrise", "sunset", "moonlight", "starlight",
"candlelight", "natural lighting", "golden hour"},
"artificial": {"neon", "fluorescent", "LED", "lamp", "candle", "torch"},
"style": {"cinematic lighting", "soft lighting", "hard lighting",
"volumetric", "rim lighting", "backlighting", "sidelight",
"god rays", "crepuscular", "studio lighting", "flash",
"warm lighting", "cool lighting", "dramatic lighting",
"moody lighting", "dark lighting", "spotlight"},
}
def analyze_background(img_arr: np.ndarray, wd14_tags: list[str]) -> list[str]:
"""Analyze background/atmosphere from image + WD14 tags.
Args:
img_arr: RGB image as numpy array (H, W, 3 uint8).
wd14_tags: General tags from WD14 ensemble.
Returns:
List of Danbooru-style background/atmosphere tags.
"""
tags: set[str] = set()
# --- Pass-through WD14 tags that match background keywords ---
wd14_lower = {t.lower() for t in wd14_tags}
all_keywords = set()
for kw_set in (_BG_LOCATION_KEYWORDS, _BG_ATMOSPHERE_KEYWORDS, _BG_LIGHTING_KEYWORDS):
for keywords in kw_set.values():
all_keywords.update(keywords)
for tag in wd14_tags:
tag_lower = tag.lower().replace("_", " ")
for kw in all_keywords:
if kw in tag_lower:
tags.add(tag)
break
# --- Lightweight color analysis on border region ---
if img_arr is not None and img_arr.size > 0 and img_arr.ndim == 3:
color_tags = _analyze_border_color(img_arr)
tags.update(color_tags)
return sorted(tags) if tags else []
def _analyze_border_color(img_arr: np.ndarray) -> set[str]:
"""Analyze border region for color temperature and lighting cues."""
tags: set[str] = set()
h, w = img_arr.shape[:2]
if h < 10 or w < 10:
return tags
border_size = min(h, w) // 16
border_size = max(border_size, 5)
# Sample border pixels (top, bottom, left, right)
border_px = np.concatenate([
img_arr[:border_size, :, :].reshape(-1, 3),
img_arr[-border_size:, :, :].reshape(-1, 3),
img_arr[:, :border_size, :].reshape(-1, 3),
img_arr[:, -border_size:, :].reshape(-1, 3),
], axis=0)
if len(border_px) == 0:
return tags
# Subsample for speed
if len(border_px) > 500:
idx = np.random.choice(len(border_px), 500, replace=False)
border_px = border_px[idx]
r_mean = float(np.mean(border_px[:, 0]))
g_mean = float(np.mean(border_px[:, 1]))
b_mean = float(np.mean(border_px[:, 2]))
# Color temperature
warm = (r_mean + g_mean) / 2
cool = b_mean
brightness = (r_mean + g_mean + b_mean) / 3
if warm > cool * 1.10:
tags.add("warm lighting")
elif cool > warm * 1.10:
tags.add("cool lighting")
# Brightness
if brightness < 40:
tags.add("dark atmosphere")
elif brightness > 220:
tags.add("bright")
# Saturation (pastel vs vivid)
r_std = float(np.std(border_px[:, 0]))
g_std = float(np.std(border_px[:, 1]))
b_std = float(np.std(border_px[:, 2]))
sat = (r_std + g_std + b_std) / 3
if sat > 60:
tags.add("vibrant colors")
elif sat < 20:
tags.add("pastel colors")
# Indoor/outdoor heuristic
border_std = float(np.mean(np.std(border_px, axis=0)))
if border_std < 30:
# Uniform border → likely sky (outdoors)
if brightness > 150:
tags.add("outdoors")
elif border_std > 50:
tags.add("indoors")
return tags