This commit is contained in:
2026-05-14 22:12:40 +08:00
parent 334845d732
commit 3ba8ba80b3
6 changed files with 252 additions and 69 deletions
+78 -5
View File
@@ -87,14 +87,87 @@ def run_sam_prompt(
box=box_arg,
multimask_output=True,
)
# masks: C x H x W
best = int(np.argmax(scores))
m = masks[best]
if m.dtype != np.bool_:
m = m > 0.5
m = _pick_best_sam_mask(masks, scores, pc, pl)
m = _refine_sam_instance_mask(m)
return m
def _mask_contains_foreground_points(
mask_bool: np.ndarray,
point_coords: np.ndarray,
point_labels: np.ndarray,
) -> bool:
for (x, y), lab in zip(point_coords, point_labels):
if int(lab) != 1:
continue
xi = int(round(float(x)))
yi = int(round(float(y)))
if yi < 0 or xi < 0 or yi >= mask_bool.shape[0] or xi >= mask_bool.shape[1]:
return False
if not bool(mask_bool[yi, xi]):
return False
return True
def _to_bool_mask(m: np.ndarray) -> np.ndarray:
if m.dtype == np.bool_:
return m
return m > 0.5
def _pick_best_sam_mask(
masks: np.ndarray,
scores: np.ndarray,
point_coords: np.ndarray,
point_labels: np.ndarray,
) -> np.ndarray:
"""
多掩膜时:优先覆盖全部前景点;在高分候选中取面积最大,减少“只切到局部/缺边”的概率。
"""
n = int(masks.shape[0])
valid_idx: List[int] = []
for i in range(n):
mb = _to_bool_mask(masks[i])
if _mask_contains_foreground_points(mb, point_coords, point_labels):
valid_idx.append(i)
if not valid_idx:
best = int(np.argmax(scores))
return _to_bool_mask(masks[best])
smax = max(float(scores[j]) for j in valid_idx)
near = [j for j in valid_idx if float(scores[j]) >= smax - 0.12]
best_area = -1
best_j = near[0]
for j in near:
mb = _to_bool_mask(masks[j])
area = int(np.count_nonzero(mb))
if area > best_area:
best_area = area
best_j = j
return _to_bool_mask(masks[best_j])
def _refine_sam_instance_mask(mask_bool: np.ndarray) -> np.ndarray:
"""填内部孔洞 + 轻微闭运算,使复杂物体轮廓更完整。"""
m = np.asarray(mask_bool, dtype=bool)
try:
from scipy import ndimage as ndi # type: ignore[import]
m = ndi.binary_fill_holes(m)
m = ndi.binary_closing(m, structure=np.ones((3, 3), dtype=bool))
return m
except Exception:
pass
try:
import cv2 # type: ignore[import]
u8 = (m.astype(np.uint8) * 255)
k = np.ones((3, 3), np.uint8)
u8 = cv2.morphologyEx(u8, cv2.MORPH_CLOSE, k)
return u8 > 127
except Exception:
return m
def expand_mask(mask_bool: np.ndarray, expand_px: int = 0) -> np.ndarray:
"""
轻微扩大二值 mask(用于消除边界过软/过细带来的“漏边”观感)。