From 4c90154bca35fb8b4edae6bc5c90fca53ee33875 Mon Sep 17 00:00:00 2001 From: Peter Schroedl Date: Mon, 25 Nov 2024 07:19:11 -0800 Subject: [PATCH] WIP add mask output --- nodes.py | 36 +++++++++++++++++++++++------------ sam2/sam2_camera_predictor.py | 6 ++++++ 2 files changed, 30 insertions(+), 12 deletions(-) diff --git a/nodes.py b/nodes.py index a420ede..979bbc1 100644 --- a/nodes.py +++ b/nodes.py @@ -2,6 +2,7 @@ import torch import os import numpy as np import logging +import json from .sam2.sam2_camera_predictor import SAM2CameraPredictor from comfy.utils import load_torch_file @@ -139,8 +140,8 @@ class Sam2RealtimeSegmentation: }, } - RETURN_NAMES = ("PROCESSED_IMAGES",) - RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("PROCESSED_IMAGES","MASK",) + RETURN_TYPES = ("IMAGE", "IMAGE",) FUNCTION = "segment_images" CATEGORY = "SAM2-Realtime" @@ -166,10 +167,10 @@ class Sam2RealtimeSegmentation: model.to(device) processed_frames = [] - + mask_list = [] # The `model` variable is now ready and equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor if self.predictor is None: - self.predictor = model + self.predictor = model def process_frame(frame, frame_idx): with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16): @@ -187,27 +188,38 @@ class Sam2RealtimeSegmentation: out_obj_ids, out_mask_logits = self.predictor.track(frame) if out_mask_logits.shape[0] > 0: + # Ensure out_mask_logits is on the same device + out_mask_logits = out_mask_logits.to(device) mask = (out_mask_logits[0, 0] > 0.5).byte() mask = torch.nn.functional.interpolate( - mask.unsqueeze(0).unsqueeze(0), + mask.unsqueeze(0).unsqueeze(0).float(), size=(frame.shape[0], frame.shape[1]), mode='nearest' - ).squeeze(0).squeeze(0) + ).squeeze(0).squeeze(0).byte().to(device) # Move the interpolated mask to the correct device else: - mask = torch.zeros((frame.shape[0], frame.shape[1]), device=device, dtype=torch.uint8) + mask = torch.ones((frame.shape[0], frame.shape[1]), device=device, dtype=torch.uint8) - mask_colored = torch.stack([mask] * 3, dim=2) # Create 3-channel mask - overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3) + # Ensure frame is on the same device + frame = frame.to(device) + + mask_colored = torch.stack([mask] * 3, dim=2).to(device) # Create 3-channel mask and move to device + + overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3).to(device) processed_frames.append(overlayed_frame) + constructed_mask = torch.add(frame * 0.1, mask_colored * 0.9).to(device) + mask_list.append(constructed_mask) + + # Avoid keeping all frames in memory for frame_idx, img in enumerate(images): process_frame(img, frame_idx) - if frame_idx % 10 == 0: - torch.cuda.empty_cache() + # if frame_idx % 10 == 0: + # torch.cuda.empty_cache() + stacked_masks = torch.stack(mask_list, dim=0) stacked_frames = torch.stack(processed_frames, dim=0) - return (stacked_frames,) + return (stacked_frames, stacked_masks) NODE_CLASS_MAPPINGS = { "DownloadAndLoadSAM2RealtimeModel": DownloadAndLoadSAM2RealtimeModel, diff --git a/sam2/sam2_camera_predictor.py b/sam2/sam2_camera_predictor.py index 36dd2e6..02da745 100644 --- a/sam2/sam2_camera_predictor.py +++ b/sam2/sam2_camera_predictor.py @@ -10,6 +10,12 @@ import torch from tqdm import tqdm +import sys +import os + +# To import from local sam2 +sys.path.append(os.path.dirname(os.path.abspath(__file__))) + from sam2.modeling.sam2_base import NO_OBJ_SCORE, SAM2Base from sam2.utils.misc import concat_points, fill_holes_in_mask_scores, load_video_frames import numpy as np