Cesium2 / src /vision.py
MORPH-AI
feat: dynamic MoE expansion, multi-head CoT, plugin architecture, improved MoD
82f262a
Raw
History Blame Contribute Delete
7.91 kB
"""
VisionAnalyzer - multimodal visual processing for MORPH-AI.
Lazy-loads a Vision Transformer (ViT) + object-detection model when
transformers provides them; otherwise falls back to pure pixel statistics so
visual facts (dominant colors, brightness, edge density, saliency regions)
are still produced with zero model downloads.
Output is a structured ImageFacts object that flows into the FSM VISION
state, feeds the RAG context, and is cross-examined by the verifier.
"""
import io
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
import torch
@dataclass
class ImageFacts:
width: int = 0
height: int = 0
dominant_colors: List[tuple] = field(default_factory=list)
brightness: float = 0.0
edge_density: float = 0.0
saliency_regions: List[Dict[str, Any]] = field(default_factory=list)
objects: List[Dict[str, Any]] = field(default_factory=list)
caption: str = ""
embedding: Optional[torch.Tensor] = None # (patches+1, dim) or None
def to_text(self) -> str:
lines = [f"image {self.width}x{self.height}", f"brightness {self.brightness:.2f}"]
if self.dominant_colors:
lines.append("colors: " + ", ".join(
f"#{r:02x}{g:02x}{b:02x}" for r, g, b in self.dominant_colors[:4]
))
if self.objects:
lines.append("objects: " + ", ".join(
f"{o.get('label', 'obj')} ({o.get('conf', 0):.2f})" for o in self.objects
))
if self.saliency_regions:
lines.append("regions: " + ", ".join(
f"{r['x']},{r['y']}" for r in self.saliency_regions[:6]
))
if self.caption:
lines.append(f"caption: {self.caption}")
return " | ".join(lines)
def to_dict(self) -> Dict[str, Any]:
return {
"width": self.width,
"height": self.height,
"dominant_colors": [list(c) for c in self.dominant_colors],
"brightness": self.brightness,
"edge_density": self.edge_density,
"saliency_regions": self.saliency_regions,
"objects": self.objects,
"caption": self.caption,
}
def _load_pil():
try:
from PIL import Image
return Image
except ImportError:
return None
class VisionAnalyzer:
def __init__(self, device: Optional[str] = None, use_vit: bool = True, use_detector: bool = True):
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
self.vit = None
self.processor = None
self.detector = None
self.det_processor = None
self.use_vit = use_vit
self.use_detector = use_detector
self._load_models()
def _load_models(self):
try:
from transformers import (
AutoImageProcessor,
AutoModelForObjectDetection,
ViTModel,
)
if self.use_vit:
self.vit = ViTModel.from_pretrained("google/vit-base-patch16-224-in21k")
self.vit = self.vit.to(self.device).eval()
if self.use_detector:
self.det_processor = AutoImageProcessor.from_pretrained(
"hustvl/yolos-small"
)
self.detector = AutoModelForObjectDetection.from_pretrained(
"hustvl/yolos-small"
)
self.detector = self.detector.to(self.device).eval()
except Exception:
self.vit = None
self.detector = None
self.processor = None
def load_image(self, source) -> Any:
"""Accept a path, file-like, or bytes; returns PIL Image or None."""
Image = _load_pil()
if Image is None:
return None
try:
if isinstance(source, (str,)):
return Image.open(source).convert("RGB")
if isinstance(source, bytes):
return Image.open(io.BytesIO(source)).convert("RGB")
if hasattr(source, "read"):
return Image.open(source).convert("RGB")
return source
except Exception:
return None
def _pixel_facts(self, img) -> ImageFacts:
Image = _load_pil()
facts = ImageFacts(width=img.width, height=img.height)
small = img.resize((32, 32))
px = list(small.getdata())
n = len(px)
r_sum = g_sum = b_sum = 0
color_hist: Dict[tuple, int] = {}
for r, g, b in px:
r_sum += r
g_sum += g
b_sum += b
key = (r // 32 * 32, g // 32 * 32, b // 32 * 32)
color_hist[key] = color_hist.get(key, 0) + 1
facts.brightness = (r_sum + g_sum + b_sum) / (3.0 * n) / 255.0
facts.dominant_colors = [
(r + 16, g + 16, b + 16) for (r, g, b), _ in
sorted(color_hist.items(), key=lambda kv: -kv[1])[:4]
]
# saliency regions: brightest / highest-variance 8x8 cells
import statistics
grid = small.resize((16, 16))
gx = list(grid.getdata())
variances = []
for i in range(16):
for j in range(16):
idx = i * 16 + j
r, g, b = gx[idx][:3]
vals = [r, g, b]
variances.append(((i * 16, j * 16), statistics.pstdev(vals)))
variances.sort(key=lambda kv: -kv[1])
facts.saliency_regions = [
{"x": x, "y": y, "score": round(v, 3)}
for (x, y), v in variances[:6]
]
# edge density via PIL edge detection
try:
import ImageFilter
except ImportError:
from PIL import ImageFilter
edges = small.convert("L").filter(ImageFilter.FIND_EDGES)
epx = list(edges.getdata())
facts.edge_density = sum(1 for v in epx if v > 100) / len(epx)
return facts
def analyze(self, source) -> ImageFacts:
img = self.load_image(source)
if img is None:
raise ValueError("Could not load image")
facts = self._pixel_facts(img)
# optional real ViT embedding
if self.vit is not None:
try:
from transformers import AutoImageProcessor
if self.processor is None:
self.processor = AutoImageProcessor.from_pretrained(
"google/vit-base-patch16-224-in21k"
)
with torch.no_grad():
inputs = self.processor(images=img, return_tensors="pt").to(self.device)
out = self.vit(**inputs)
facts.embedding = out.last_hidden_state # (1, patches+1, dim)
except Exception:
facts.embedding = None
# optional object detection
if self.detector is not None:
try:
with torch.no_grad():
det = self.det_processor(images=img, return_tensors="pt").to(self.device)
outputs = self.detector(**det)
target_sizes = torch.tensor([[img.height, img.width]])
results = self.det_processor.post_process_object_detection(
outputs, threshold=0.5, target_sizes=target_sizes
)[0]
for score, label, box in zip(
results["scores"].tolist(),
results["labels"].tolist(),
results["boxes"].tolist(),
):
label_str = self.detector.config.id2label.get(label, "obj")
facts.objects.append({
"label": label_str,
"conf": score,
"box": [round(b, 1) for b in box],
})
except Exception:
pass
return facts