139 lines
5.8 KiB
Python
139 lines
5.8 KiB
Python
"""
|
|
Flower-instance counting via SAM3 (facebook/sam3), prompted with the plain-
|
|
English concept "flower". See pipeline/sam3_client.py (and config.py's SAM3
|
|
section) for why this runs as a one-shot subprocess into a separate Python
|
|
3.10 environment rather than an in-process model call.
|
|
|
|
Unlike the earlier approach (SAM1 in automatic "segment everything" mode,
|
|
then guessing which proposals were flowers from size/position heuristics
|
|
and excluding boxes YOLO-World identified as vase/ribbon), SAM3's concept
|
|
prompting does the semantic part itself: prompted with "flower", it simply
|
|
never proposes the vase or any ribbon/bow in the first place. What's left
|
|
here is only: a light sanity filter (the instance must actually overlap the
|
|
already-known foreground), then color-clustering the survivors as a rough
|
|
proxy for distinct flower "kinds" -- there's still no trained species
|
|
classifier, just an assumption that different flower types usually differ
|
|
in color.
|
|
"""
|
|
|
|
import logging
|
|
|
|
import cv2
|
|
import numpy as np
|
|
|
|
import config
|
|
from pipeline import sam3_client
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def union_mask(instance_masks):
|
|
"""OR-combines every per-flower instance mask into one -- "all the
|
|
flower material, regardless of which bloom it belongs to". Used to crop
|
|
just the flowers (excluding vase/ribbon/background) out of a photo for
|
|
the CLIP flower-similarity check. Returns None if there are no
|
|
instances to combine."""
|
|
if not instance_masks:
|
|
return None
|
|
union = instance_masks[0].copy()
|
|
for m in instance_masks[1:]:
|
|
union |= m
|
|
return union
|
|
|
|
|
|
def count_flowers(bgr, mask, workdir, image_key="input", instances_raw=None):
|
|
"""instances_raw, if given, reuses SAM3 results already fetched by the
|
|
caller (engine.py batches the "flower" job for the upload together with
|
|
any "vase" jobs into a single sam3_client.run_jobs() call so the model
|
|
only loads once per request); otherwise fetches them here standalone."""
|
|
fg_area = int((mask > 0).sum())
|
|
if fg_area == 0:
|
|
return {"total_count": 0, "clusters": [], "_instance_masks": [], "_cluster_of_instance": []}
|
|
|
|
if instances_raw is None:
|
|
result = sam3_client.run_jobs(
|
|
images={image_key: bgr},
|
|
jobs=[{"image": image_key, "prompt": config.SAM3_FLOWER_PROMPT,
|
|
"threshold": config.SAM3_FLOWER_THRESHOLD}],
|
|
workdir=workdir,
|
|
)
|
|
instances_raw = result.get((image_key, config.SAM3_FLOWER_PROMPT), [])
|
|
|
|
# Light sanity filter only -- SAM3 already did the semantic work of
|
|
# "is this a flower", this just guards against a stray instance
|
|
# entirely outside the known foreground (shouldn't happen against an
|
|
# already background-removed image, but costs nothing to check).
|
|
instance_masks = [
|
|
inst["mask"] for inst in instances_raw
|
|
if np.logical_and(inst["mask"], mask > 0).any()
|
|
]
|
|
|
|
# Color-cluster the surviving instances as a proxy for distinct flower
|
|
# "kinds" -- same idea as the color-family grid elsewhere in this app
|
|
# (pipeline/color_grid.py), just applied per-instance instead of
|
|
# per-pixel-region.
|
|
avg_colors_lab = []
|
|
for m in instance_masks:
|
|
pixels = bgr[m].reshape(-1, 1, 3).astype(np.uint8)
|
|
lab = cv2.cvtColor(pixels, cv2.COLOR_BGR2LAB).reshape(-1, 3)
|
|
avg_colors_lab.append(lab.mean(axis=0))
|
|
|
|
clusters = []
|
|
cluster_of_instance = []
|
|
if avg_colors_lab:
|
|
pts = np.array(avg_colors_lab, dtype=np.float32)
|
|
k = min(config.SAM_MAX_KIND_CLUSTERS, len(pts))
|
|
if k <= 1:
|
|
labels = np.zeros(len(pts), dtype=int)
|
|
centers = pts
|
|
else:
|
|
criteria = (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 30, 0.5)
|
|
_, labels, centers = cv2.kmeans(pts, k, None, criteria, 5, cv2.KMEANS_PP_CENTERS)
|
|
labels = labels.flatten()
|
|
cluster_of_instance = labels.tolist()
|
|
for c in range(len(centers)):
|
|
count = int((labels == c).sum())
|
|
if count == 0:
|
|
continue
|
|
lab_center = centers[c].reshape(1, 1, 3).astype(np.uint8)
|
|
bgr_center = cv2.cvtColor(lab_center, cv2.COLOR_LAB2BGR)[0, 0]
|
|
clusters.append({
|
|
"count": count,
|
|
"color_rgb": [int(bgr_center[2]), int(bgr_center[1]), int(bgr_center[0])],
|
|
})
|
|
clusters.sort(key=lambda c: c["count"], reverse=True)
|
|
|
|
return {
|
|
"total_count": len(instance_masks),
|
|
"clusters": clusters,
|
|
"_instance_masks": instance_masks,
|
|
"_cluster_of_instance": cluster_of_instance,
|
|
}
|
|
|
|
|
|
_PALETTE = [
|
|
(66, 133, 244), (219, 68, 55), (244, 180, 0), (15, 157, 88),
|
|
(171, 71, 188), (255, 112, 67), (0, 172, 193), (158, 157, 36),
|
|
]
|
|
|
|
|
|
def render_instances(bgr, count_data):
|
|
"""Overlays a translucent, distinctly-colored fill per detected flower
|
|
instance (color = its cluster/"kind") plus its index, so the count is
|
|
visually verifiable rather than just a bare number."""
|
|
overlay = bgr.copy()
|
|
instance_masks = count_data.get("_instance_masks", [])
|
|
cluster_of_instance = count_data.get("_cluster_of_instance", [])
|
|
for idx, m in enumerate(instance_masks):
|
|
cluster_id = cluster_of_instance[idx] if idx < len(cluster_of_instance) else idx
|
|
color = np.array(_PALETTE[cluster_id % len(_PALETTE)], dtype=np.float32)
|
|
overlay[m] = (color * 0.55 + overlay[m].astype(np.float32) * 0.45).astype(np.uint8)
|
|
ys, xs = np.nonzero(m)
|
|
if len(xs):
|
|
cx, cy = int(xs.mean()), int(ys.mean())
|
|
cv2.putText(overlay, str(idx + 1), (cx - 8, cy + 6),
|
|
cv2.FONT_HERSHEY_SIMPLEX, 0.6, (255, 255, 255), 2, cv2.LINE_AA)
|
|
cv2.putText(overlay, str(idx + 1), (cx - 8, cy + 6),
|
|
cv2.FONT_HERSHEY_SIMPLEX, 0.6, (30, 30, 30), 1, cv2.LINE_AA)
|
|
return overlay
|