Skip to content

Commit 8ea1034

Browse files
committed
Changes made
1 parent 4c03585 commit 8ea1034

5 files changed

Lines changed: 204 additions & 259 deletions

File tree

nodes/mask_preview.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -97,8 +97,8 @@ def preview(self, image, mask, display_mode, overlay_color_r, overlay_color_g,
9797
overlay = color.view(1, 1, 3).expand(H, W, 3)
9898
frame = frame * (1 - edge_3d) + overlay * edge_3d
9999

100-
# Edge overlay for non-edge modes
101-
if display_mode != "edge_highlight" and edge_width > 0:
100+
# Edge overlay for non-edge modes (skip side_by_side — frame width differs)
101+
if display_mode not in ("edge_highlight", "side_by_side") and edge_width > 0:
102102
edge = self._detect_edge(mi, edge_width)
103103
edge_3d = edge.unsqueeze(-1)
104104
overlay_edge = color.view(1, 1, 3).expand(H, W, 3)

nodes/sam_mask_generator.py

Lines changed: 20 additions & 140 deletions
Original file line numberDiff line numberDiff line change
@@ -9,11 +9,14 @@
99
import json
1010
import gc
1111

12-
try:
13-
import cv2
14-
HAS_CV2 = True
15-
except ImportError:
16-
HAS_CV2 = False
12+
from .utils import (
13+
HAS_CV2,
14+
get_sam_predictor,
15+
augment_prompts_from_mask,
16+
mask_to_sam_logits,
17+
parse_points_json,
18+
parse_bbox_input,
19+
)
1720

1821

1922
class SAMMaskGeneratorMEC:
@@ -127,42 +130,20 @@ def _run_inference(self, model, model_type, image, points_json, bbox_json,
127130
img_np = (img_tensor.cpu().numpy() * 255).astype(np.uint8)
128131
H, W = img_np.shape[:2]
129132

130-
# Parse points
131-
try:
132-
points_list = json.loads(points_json) if isinstance(points_json, str) else points_json
133-
except json.JSONDecodeError:
134-
points_list = []
135-
136-
point_coords = None
137-
point_labels = None
133+
# Parse points (shared utility)
134+
points_list = parse_points_json(points_json)
135+
point_coords, point_labels = None, None
138136
if points_list:
139137
coords = [[float(p["x"]), float(p["y"])] for p in points_list]
140138
labels = [int(p.get("label", 1)) for p in points_list]
141139
point_coords = np.array(coords, dtype=np.float32)
142140
point_labels = np.array(labels, dtype=np.int32)
143141

144-
# Parse bbox
145-
box_np = None
146-
if bbox_input is not None:
147-
# From BBOX node: [x, y, w, h] → [x1, y1, x2, y2]
148-
bx, by, bw, bh = bbox_input
149-
box_np = np.array([bx, by, bx + bw, by + bh], dtype=np.float32)
150-
elif bbox_json and bbox_json.strip():
151-
try:
152-
bdata = json.loads(bbox_json)
153-
if isinstance(bdata, list) and len(bdata) == 4:
154-
box_np = np.array(bdata, dtype=np.float32)
155-
elif isinstance(bdata, dict):
156-
bx = float(bdata.get("x", 0))
157-
by = float(bdata.get("y", 0))
158-
bw = float(bdata.get("w", bdata.get("width", 0)))
159-
bh = float(bdata.get("h", bdata.get("height", 0)))
160-
box_np = np.array([bx, by, bx + bw, by + bh], dtype=np.float32)
161-
except json.JSONDecodeError:
162-
pass
163-
164-
# ── Build predictor ────────────────────────────────────────────
165-
predictor = self._get_predictor(model, model_type, img_np)
142+
# Parse bbox (shared utility)
143+
box_np = parse_bbox_input(bbox_json, bbox_input)
144+
145+
# ── Build predictor (shared utility) ───────────────────────────
146+
predictor = get_sam_predictor(model, model_type, img_np)
166147

167148
if predictor is None:
168149
empty = torch.zeros(1, H, W, dtype=torch.float32)
@@ -188,17 +169,17 @@ def _run_inference(self, model, model_type, image, points_json, bbox_json,
188169
iter_labels = point_labels
189170
iter_box = box_np
190171

191-
# Augment prompts from previous iteration
172+
# Augment prompts from previous iteration (shared utility)
192173
if iteration > 0 and current_mask is not None:
193-
iter_coords, iter_labels, iter_box = self._augment_prompts(
174+
iter_coords, iter_labels, iter_box = augment_prompts_from_mask(
194175
current_mask, point_coords, point_labels, box_np,
195-
auto_negative_points, H, W,
176+
H, W, auto_negative=auto_negative_points,
196177
)
197178

198179
# Prepare mask input (logits from previous pass)
199180
mask_input = None
200181
if iteration > 0 and current_mask is not None:
201-
mask_input = self._mask_to_logits(current_mask)
182+
mask_input = mask_to_sam_logits(current_mask)
202183

203184
# Run SAM
204185
try:
@@ -296,107 +277,6 @@ def _run_inference(self, model, model_type, image, points_json, bbox_json,
296277

297278
return (selected_mask, all_masks_t, det_bbox, selected_score, info)
298279

299-
# ── Predictor factory ─────────────────────────────────────────────
300-
301-
def _get_predictor(self, model, model_type, img_np):
302-
"""Get the correct predictor for the model type and set image."""
303-
predictor = None
304-
305-
if model_type in ("sam2", "sam2.1"):
306-
try:
307-
from sam2.sam2_image_predictor import SAM2ImagePredictor
308-
predictor = SAM2ImagePredictor(model)
309-
except Exception:
310-
pass
311-
312-
elif model_type == "sam3":
313-
try:
314-
from sam3.predictor import SAM3Predictor
315-
predictor = SAM3Predictor(model)
316-
except Exception:
317-
pass
318-
319-
else:
320-
try:
321-
from segment_anything import SamPredictor
322-
predictor = SamPredictor(model)
323-
except Exception:
324-
pass
325-
326-
if predictor is not None:
327-
try:
328-
predictor.set_image(img_np)
329-
except Exception:
330-
return None
331-
332-
return predictor
333-
334-
# ── Iterative prompt augmentation ─────────────────────────────────
335-
336-
def _augment_prompts(self, mask, orig_coords, orig_labels, orig_box,
337-
auto_neg, H, W):
338-
"""Generate augmented prompts from the previous mask pass."""
339-
coords_list = []
340-
labels_list = []
341-
342-
if orig_coords is not None:
343-
coords_list.append(orig_coords)
344-
labels_list.append(orig_labels)
345-
346-
if not HAS_CV2:
347-
c = np.concatenate(coords_list) if coords_list else None
348-
l = np.concatenate(labels_list) if labels_list else None
349-
return c, l, orig_box
350-
351-
binary = (mask > 0.5).astype(np.uint8)
352-
353-
# Interior positive points
354-
kern = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))
355-
interior = cv2.erode(binary, kern, iterations=1)
356-
pts = np.argwhere(interior > 0)
357-
if len(pts) > 0:
358-
n = min(3, len(pts))
359-
idx = np.linspace(0, len(pts) - 1, n, dtype=int)
360-
coords_list.append(pts[idx][:, ::-1].astype(np.float32))
361-
labels_list.append(np.ones(n, dtype=np.int32))
362-
363-
# Exterior negative points
364-
if auto_neg:
365-
dilated = cv2.dilate(binary, kern, iterations=1)
366-
exterior = dilated - binary
367-
pts = np.argwhere(exterior > 0)
368-
if len(pts) > 0:
369-
n = min(2, len(pts))
370-
idx = np.linspace(0, len(pts) - 1, n, dtype=int)
371-
coords_list.append(pts[idx][:, ::-1].astype(np.float32))
372-
labels_list.append(np.zeros(n, dtype=np.int32))
373-
374-
all_coords = np.concatenate(coords_list).astype(np.float32) if coords_list else None
375-
all_labels = np.concatenate(labels_list).astype(np.int32) if labels_list else None
376-
377-
# Derive tighter bbox
378-
ys, xs = np.where(binary > 0)
379-
if len(xs) > 0:
380-
pad = max(5, int(min(H, W) * 0.02))
381-
box = np.array([
382-
max(0, xs.min() - pad), max(0, ys.min() - pad),
383-
min(W, xs.max() + pad), min(H, ys.max() + pad),
384-
], dtype=np.float32)
385-
else:
386-
box = orig_box
387-
388-
return all_coords, all_labels, box
389-
390-
@staticmethod
391-
def _mask_to_logits(mask, target_size=256):
392-
"""Convert float mask → SAM logit space (inverse sigmoid) at 256x256."""
393-
m = np.clip(mask, 1e-6, 1 - 1e-6)
394-
logits = np.log(m / (1 - m))
395-
if HAS_CV2:
396-
logits = cv2.resize(logits, (target_size, target_size),
397-
interpolation=cv2.INTER_LINEAR)
398-
return logits[np.newaxis, :, :] # (1, 256, 256)
399-
400280
@staticmethod
401281
def _fallback_predict(model, img_np, point_coords, point_labels, box_np, multimask):
402282
"""Generic fallback using model's forward pass directly."""

nodes/sam_vitmatte_pipeline.py

Lines changed: 12 additions & 110 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,11 @@
3636
remove_small_regions,
3737
mask_to_bbox,
3838
make_split_preview,
39+
augment_prompts_from_mask,
40+
mask_to_sam_logits,
41+
parse_points_json,
42+
points_to_arrays,
43+
parse_bbox_input,
3944
)
4045

4146
try:
@@ -147,9 +152,9 @@ def execute(self, sam_model, image, points_json, bbox_json,
147152
img_np = (img_tensor.cpu().numpy() * 255).astype(np.uint8)
148153
H, W = img_np.shape[:2]
149154

150-
# Parse prompts
151-
points_list = self._parse_points(points_json)
152-
box_np = self._parse_bbox(bbox_json, bbox)
155+
# Parse prompts (shared utilities)
156+
points_list = parse_points_json(points_json)
157+
box_np = parse_bbox_input(bbox_json, bbox)
153158

154159
# ── Stage 1: SAM coarse mask (with iterative refinement) ──────
155160
if offload and hasattr(model, "to"):
@@ -194,7 +199,7 @@ def execute(self, sam_model, image, points_json, bbox_json,
194199
# ── Outputs ───────────────────────────────────────────────────
195200
edge_mask = torch.abs(refined_mask - (coarse_mask > 0.5).float())
196201
det_bbox = mask_to_bbox(refined_mask, W, H)
197-
preview = make_split_preview(img_tensor, refined_mask, coarse_mask)
202+
preview = make_split_preview(img_tensor, coarse_mask, refined_mask)
198203

199204
info = json.dumps({
200205
"model_type": model_type,
@@ -232,7 +237,7 @@ def _iterative_sam(self, model, model_type, img_np, points_list,
232237
if predictor is None:
233238
return torch.zeros((H, W), dtype=torch.float32), 0.0
234239

235-
point_coords, point_labels = self._points_to_arrays(points_list)
240+
point_coords, point_labels = points_to_arrays(points_list)
236241

237242
# Use existing mask as starting point if provided
238243
current_mask = None
@@ -251,7 +256,7 @@ def _iterative_sam(self, model, model_type, img_np, points_list,
251256
iter_box = box_np
252257

253258
if iteration > 0 and current_mask is not None:
254-
aug_coords, aug_labels, aug_box = self._augment_prompts_from_mask(
259+
aug_coords, aug_labels, aug_box = augment_prompts_from_mask(
255260
current_mask, point_coords, point_labels, box_np, H, W
256261
)
257262
iter_coords = aug_coords
@@ -263,14 +268,7 @@ def _iterative_sam(self, model, model_type, img_np, points_list,
263268
try:
264269
mask_input = None
265270
if iteration > 0 and current_mask is not None:
266-
m_logit = np.clip(current_mask, 1e-6, 1 - 1e-6)
267-
m_logit = np.log(m_logit / (1 - m_logit))
268-
if HAS_CV2:
269-
mask_input = cv2.resize(m_logit, (256, 256),
270-
interpolation=cv2.INTER_LINEAR)
271-
else:
272-
mask_input = m_logit
273-
mask_input = mask_input[np.newaxis, :, :] # (1, 256, 256)
271+
mask_input = mask_to_sam_logits(current_mask)
274272

275273
masks_np, scores, _ = predictor.predict(
276274
point_coords=iter_coords,
@@ -310,63 +308,6 @@ def _iterative_sam(self, model, model_type, img_np, points_list,
310308

311309
return torch.from_numpy(current_mask), best_score
312310

313-
def _augment_prompts_from_mask(self, mask, orig_coords, orig_labels,
314-
orig_box, H, W):
315-
"""Generate additional prompts from the previous mask iteration."""
316-
coords_list = []
317-
labels_list = []
318-
319-
if orig_coords is not None:
320-
coords_list.append(orig_coords)
321-
labels_list.append(orig_labels)
322-
323-
if not HAS_CV2:
324-
c = np.concatenate(coords_list) if coords_list else None
325-
l = np.concatenate(labels_list) if labels_list else None
326-
return c, l, orig_box
327-
328-
binary = (mask > 0.5).astype(np.uint8)
329-
330-
# Eroded interior → strong positive points
331-
kern = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))
332-
interior = cv2.erode(binary, kern, iterations=1)
333-
interior_pts = np.argwhere(interior > 0) # (y, x)
334-
if len(interior_pts) > 0:
335-
n = min(3, len(interior_pts))
336-
indices = np.linspace(0, len(interior_pts) - 1, n, dtype=int)
337-
sampled = interior_pts[indices]
338-
coords_list.append(sampled[:, ::-1].astype(np.float32))
339-
labels_list.append(np.ones(n, dtype=np.int32))
340-
341-
# Dilated boundary exterior → negative points
342-
dilated = cv2.dilate(binary, kern, iterations=1)
343-
exterior = dilated - binary
344-
exterior_pts = np.argwhere(exterior > 0)
345-
if len(exterior_pts) > 0:
346-
n = min(2, len(exterior_pts))
347-
indices = np.linspace(0, len(exterior_pts) - 1, n, dtype=int)
348-
sampled = exterior_pts[indices]
349-
coords_list.append(sampled[:, ::-1].astype(np.float32))
350-
labels_list.append(np.zeros(n, dtype=np.int32))
351-
352-
all_coords = np.concatenate(coords_list).astype(np.float32) if coords_list else None
353-
all_labels = np.concatenate(labels_list).astype(np.int32) if labels_list else None
354-
355-
# Derive tighter bbox from mask
356-
ys, xs = np.where(binary > 0)
357-
if len(xs) > 0:
358-
pad = max(5, int(min(H, W) * 0.02))
359-
box = np.array([
360-
max(0, xs.min() - pad),
361-
max(0, ys.min() - pad),
362-
min(W, xs.max() + pad),
363-
min(H, ys.max() + pad),
364-
], dtype=np.float32)
365-
else:
366-
box = orig_box
367-
368-
return all_coords, all_labels, box
369-
370311
# ══════════════════════════════════════════════════════════════════
371312
# STAGE 3 – Edge-aware matting (delegates to utils)
372313
# ══════════════════════════════════════════════════════════════════
@@ -456,42 +397,3 @@ def _try_laplacian(mask_np, edge_radius, detail_pres):
456397
return torch.from_numpy(np.clip(result, 0, 1).astype(np.float32))
457398
except Exception:
458399
return None
459-
460-
# ══════════════════════════════════════════════════════════════════
461-
# Prompt parsing
462-
# ══════════════════════════════════════════════════════════════════
463-
464-
@staticmethod
465-
def _parse_points(points_json):
466-
try:
467-
return json.loads(points_json) if isinstance(points_json, str) else points_json
468-
except json.JSONDecodeError:
469-
return []
470-
471-
@staticmethod
472-
def _points_to_arrays(points_list):
473-
if not points_list:
474-
return None, None
475-
coords = [[float(p["x"]), float(p["y"])] for p in points_list]
476-
labels = [int(p.get("label", 1)) for p in points_list]
477-
return np.array(coords, dtype=np.float32), np.array(labels, dtype=np.int32)
478-
479-
@staticmethod
480-
def _parse_bbox(bbox_json, bbox_input):
481-
if bbox_input is not None:
482-
bx, by, bw, bh = bbox_input
483-
return np.array([bx, by, bx + bw, by + bh], dtype=np.float32)
484-
if bbox_json and bbox_json.strip():
485-
try:
486-
bdata = json.loads(bbox_json)
487-
if isinstance(bdata, list) and len(bdata) == 4:
488-
return np.array(bdata, dtype=np.float32)
489-
elif isinstance(bdata, dict):
490-
bx = float(bdata.get("x", 0))
491-
by = float(bdata.get("y", 0))
492-
bw = float(bdata.get("w", bdata.get("width", 0)))
493-
bh = float(bdata.get("h", bdata.get("height", 0)))
494-
return np.array([bx, by, bx + bw, by + bh], dtype=np.float32)
495-
except json.JSONDecodeError:
496-
pass
497-
return None

0 commit comments

Comments
 (0)