update
This commit is contained in:
@@ -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(用于消除边界过软/过细带来的“漏边”观感)。
|
||||
|
||||
Reference in New Issue
Block a user