update
This commit is contained in:
@@ -75,7 +75,7 @@ def segment_sam_prompt(req: SamPromptSegmentRequest) -> Dict[str, Any]:
|
|||||||
return {"success": False, "error": "SAM 未产生有效掩膜", "contour": []}
|
return {"success": False, "error": "SAM 未产生有效掩膜", "contour": []}
|
||||||
|
|
||||||
mask = expand_mask(mask, int(req.expand_px))
|
mask = expand_mask(mask, int(req.expand_px))
|
||||||
contour = mask_to_contour_xy(mask, epsilon_px=2.0)
|
contour = mask_to_contour_xy(mask, epsilon_px=3.5)
|
||||||
if len(contour) < 3:
|
if len(contour) < 3:
|
||||||
return {"success": False, "error": "轮廓点数不足", "contour": []}
|
return {"success": False, "error": "轮廓点数不足", "contour": []}
|
||||||
|
|
||||||
|
|||||||
@@ -81,13 +81,23 @@ def run_sam_prompt(
|
|||||||
b = np.asarray(box_xyxy, dtype=np.float32).reshape(4)
|
b = np.asarray(box_xyxy, dtype=np.float32).reshape(4)
|
||||||
box_arg = b
|
box_arg = b
|
||||||
|
|
||||||
masks, scores, _low = predictor.predict(
|
masks3, scores3, _low3 = predictor.predict(
|
||||||
point_coords=pc,
|
point_coords=pc,
|
||||||
point_labels=pl,
|
point_labels=pl,
|
||||||
box=box_arg,
|
box=box_arg,
|
||||||
multimask_output=True,
|
multimask_output=True,
|
||||||
)
|
)
|
||||||
m = _pick_best_sam_mask(masks, scores, pc, pl)
|
m_multi = _union_sam_masks_covering_foreground(masks3, scores3, pc, pl)
|
||||||
|
|
||||||
|
masks1, scores1, _low1 = predictor.predict(
|
||||||
|
point_coords=pc,
|
||||||
|
point_labels=pl,
|
||||||
|
box=box_arg,
|
||||||
|
multimask_output=False,
|
||||||
|
)
|
||||||
|
m_single = _to_bool_mask(masks1[0])
|
||||||
|
|
||||||
|
m = _prefer_larger_foreground_mask(m_multi, m_single, pc, pl)
|
||||||
m = _refine_sam_instance_mask(m)
|
m = _refine_sam_instance_mask(m)
|
||||||
return m
|
return m
|
||||||
|
|
||||||
@@ -96,15 +106,26 @@ def _mask_contains_foreground_points(
|
|||||||
mask_bool: np.ndarray,
|
mask_bool: np.ndarray,
|
||||||
point_coords: np.ndarray,
|
point_coords: np.ndarray,
|
||||||
point_labels: np.ndarray,
|
point_labels: np.ndarray,
|
||||||
|
*,
|
||||||
|
dilate_for_hit: int = 0,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
"""dilate_for_hit>0 时:前景点落在膨胀后的掩膜内即算覆盖(缓解边缘亚像素偏差)。"""
|
||||||
|
m = np.asarray(mask_bool, dtype=bool)
|
||||||
|
if dilate_for_hit > 0:
|
||||||
|
try:
|
||||||
|
from scipy import ndimage as ndi # type: ignore[import]
|
||||||
|
|
||||||
|
m = ndi.binary_dilation(m, structure=np.ones((3, 3), dtype=bool), iterations=dilate_for_hit)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
for (x, y), lab in zip(point_coords, point_labels):
|
for (x, y), lab in zip(point_coords, point_labels):
|
||||||
if int(lab) != 1:
|
if int(lab) != 1:
|
||||||
continue
|
continue
|
||||||
xi = int(round(float(x)))
|
xi = int(round(float(x)))
|
||||||
yi = int(round(float(y)))
|
yi = int(round(float(y)))
|
||||||
if yi < 0 or xi < 0 or yi >= mask_bool.shape[0] or xi >= mask_bool.shape[1]:
|
if yi < 0 or xi < 0 or yi >= m.shape[0] or xi >= m.shape[1]:
|
||||||
return False
|
return False
|
||||||
if not bool(mask_bool[yi, xi]):
|
if not bool(m[yi, xi]):
|
||||||
return False
|
return False
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -115,45 +136,74 @@ def _to_bool_mask(m: np.ndarray) -> np.ndarray:
|
|||||||
return m > 0.5
|
return m > 0.5
|
||||||
|
|
||||||
|
|
||||||
def _pick_best_sam_mask(
|
def _union_sam_masks_covering_foreground(
|
||||||
masks: np.ndarray,
|
masks: np.ndarray,
|
||||||
scores: np.ndarray,
|
scores: np.ndarray,
|
||||||
point_coords: np.ndarray,
|
point_coords: np.ndarray,
|
||||||
point_labels: np.ndarray,
|
point_labels: np.ndarray,
|
||||||
) -> np.ndarray:
|
) -> np.ndarray:
|
||||||
"""
|
"""
|
||||||
多掩膜时:优先覆盖全部前景点;在高分候选中取面积最大,减少“只切到局部/缺边”的概率。
|
对「包含全部前景点」的候选掩膜做逻辑并,通常比单选一片更接近完整物体(船身+棚+人物等)。
|
||||||
|
若并集过大(误吞整图)则退回「有效候选中面积最大」的一片。
|
||||||
"""
|
"""
|
||||||
n = int(masks.shape[0])
|
n = int(masks.shape[0])
|
||||||
valid_idx: List[int] = []
|
valid: List[np.ndarray] = []
|
||||||
for i in range(n):
|
for i in range(n):
|
||||||
mb = _to_bool_mask(masks[i])
|
mb = _to_bool_mask(masks[i])
|
||||||
if _mask_contains_foreground_points(mb, point_coords, point_labels):
|
if _mask_contains_foreground_points(mb, point_coords, point_labels, dilate_for_hit=2):
|
||||||
valid_idx.append(i)
|
valid.append(mb)
|
||||||
if not valid_idx:
|
if not valid:
|
||||||
best = int(np.argmax(scores))
|
best = int(np.argmax(scores))
|
||||||
return _to_bool_mask(masks[best])
|
return _to_bool_mask(masks[best])
|
||||||
smax = max(float(scores[j]) for j in valid_idx)
|
if len(valid) == 1:
|
||||||
near = [j for j in valid_idx if float(scores[j]) >= smax - 0.12]
|
return valid[0]
|
||||||
best_area = -1
|
u = np.zeros_like(valid[0], dtype=bool)
|
||||||
best_j = near[0]
|
for v in valid:
|
||||||
for j in near:
|
u |= v
|
||||||
mb = _to_bool_mask(masks[j])
|
area_u = int(np.count_nonzero(u))
|
||||||
area = int(np.count_nonzero(mb))
|
h, w = u.shape[:2]
|
||||||
if area > best_area:
|
if area_u > int(0.62 * h * w):
|
||||||
best_area = area
|
return max(valid, key=lambda x: int(np.count_nonzero(x)))
|
||||||
best_j = j
|
return u
|
||||||
return _to_bool_mask(masks[best_j])
|
|
||||||
|
|
||||||
|
def _prefer_larger_foreground_mask(
|
||||||
|
a: np.ndarray,
|
||||||
|
b: np.ndarray,
|
||||||
|
point_coords: np.ndarray,
|
||||||
|
point_labels: np.ndarray,
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""在「多掩膜并集」与「单掩膜」之间选:均含前景点时优先面积明显更大者(更接近整物体)。"""
|
||||||
|
a = np.asarray(a, dtype=bool)
|
||||||
|
b = np.asarray(b, dtype=bool)
|
||||||
|
ca = _mask_contains_foreground_points(a, point_coords, point_labels, dilate_for_hit=2)
|
||||||
|
cb = _mask_contains_foreground_points(b, point_coords, point_labels, dilate_for_hit=2)
|
||||||
|
aa = int(np.count_nonzero(a))
|
||||||
|
ab = int(np.count_nonzero(b))
|
||||||
|
|
||||||
|
if ca and cb:
|
||||||
|
if ab > aa * 1.08:
|
||||||
|
return b
|
||||||
|
if aa > ab * 1.08:
|
||||||
|
return a
|
||||||
|
return a if aa >= ab else b
|
||||||
|
if ca:
|
||||||
|
return a
|
||||||
|
if cb:
|
||||||
|
return b
|
||||||
|
return a if aa >= ab else b
|
||||||
|
|
||||||
|
|
||||||
def _refine_sam_instance_mask(mask_bool: np.ndarray) -> np.ndarray:
|
def _refine_sam_instance_mask(mask_bool: np.ndarray) -> np.ndarray:
|
||||||
"""填内部孔洞 + 轻微闭运算,使复杂物体轮廓更完整。"""
|
"""填孔 + 闭运算连接断裂,再轻微膨胀以包全外轮廓。"""
|
||||||
m = np.asarray(mask_bool, dtype=bool)
|
m = np.asarray(mask_bool, dtype=bool)
|
||||||
try:
|
try:
|
||||||
from scipy import ndimage as ndi # type: ignore[import]
|
from scipy import ndimage as ndi # type: ignore[import]
|
||||||
|
|
||||||
m = ndi.binary_fill_holes(m)
|
m = ndi.binary_fill_holes(m)
|
||||||
|
m = ndi.binary_closing(m, structure=np.ones((5, 5), dtype=bool))
|
||||||
m = ndi.binary_closing(m, structure=np.ones((3, 3), dtype=bool))
|
m = ndi.binary_closing(m, structure=np.ones((3, 3), dtype=bool))
|
||||||
|
m = ndi.binary_dilation(m, structure=np.ones((3, 3), dtype=bool), iterations=1)
|
||||||
return m
|
return m
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
@@ -161,8 +211,9 @@ def _refine_sam_instance_mask(mask_bool: np.ndarray) -> np.ndarray:
|
|||||||
import cv2 # type: ignore[import]
|
import cv2 # type: ignore[import]
|
||||||
|
|
||||||
u8 = (m.astype(np.uint8) * 255)
|
u8 = (m.astype(np.uint8) * 255)
|
||||||
k = np.ones((3, 3), np.uint8)
|
k5 = np.ones((5, 5), np.uint8)
|
||||||
u8 = cv2.morphologyEx(u8, cv2.MORPH_CLOSE, k)
|
u8 = cv2.morphologyEx(u8, cv2.MORPH_CLOSE, k5)
|
||||||
|
u8 = cv2.dilate(u8, np.ones((3, 3), np.uint8), iterations=1)
|
||||||
return u8 > 127
|
return u8 > 127
|
||||||
except Exception:
|
except Exception:
|
||||||
return m
|
return m
|
||||||
@@ -239,7 +290,7 @@ def mask_to_contour_xy(
|
|||||||
if cv2.contourArea(cnt) < 1.0:
|
if cv2.contourArea(cnt) < 1.0:
|
||||||
return []
|
return []
|
||||||
peri = cv2.arcLength(cnt, True)
|
peri = cv2.arcLength(cnt, True)
|
||||||
eps = max(epsilon_px, 0.001 * peri)
|
eps = max(epsilon_px, 0.0025 * peri)
|
||||||
approx = cv2.approxPolyDP(cnt, eps, True)
|
approx = cv2.approxPolyDP(cnt, eps, True)
|
||||||
out: List[List[float]] = []
|
out: List[List[float]] = []
|
||||||
for p in approx:
|
for p in approx:
|
||||||
|
|||||||
Reference in New Issue
Block a user