wip fixup automasking, show coord point
This commit is contained in:
@@ -8,6 +8,8 @@ import ast
|
||||
|
||||
import sys
|
||||
|
||||
import cv2
|
||||
|
||||
# Add the directory containing 'sam2_realtime' to sys.path
|
||||
current_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
sam2_realtime_path = os.path.join(current_directory) # Adjust the relative path
|
||||
@@ -83,7 +85,7 @@ class DownloadAndLoadSAM2RealtimeModel:
|
||||
config_dir = os.path.join(script_directory, "sam2_configs")
|
||||
|
||||
# Code ripped out of sam2.build_sam.build_sam2_camera_predictor to appease Hydra
|
||||
model_cfg = "sam2_hiera_t.yaml" #TODO: remove hardcoded config and path
|
||||
model_cfg = "sam2_hiera_t.yaml" #TODO(pschroedl): remove hardcoded config and path
|
||||
with initialize_config_dir(config_dir=config_dir, version_base=None):
|
||||
cfg = compose(config_name=model_cfg)
|
||||
|
||||
@@ -148,6 +150,8 @@ class Sam2RealtimeSegmentation:
|
||||
# "bboxes": ("BBOX", ),
|
||||
# "individual_objects": ("BOOLEAN", {"default": False}),
|
||||
# "mask": ("MASK", ),
|
||||
"threshold": ("FLOAT", {"forceInput": True}),
|
||||
"show_point": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -160,6 +164,40 @@ class Sam2RealtimeSegmentation:
|
||||
self.predictor = None
|
||||
self.if_init = False
|
||||
|
||||
|
||||
def _process_mask(self, mask: np.ndarray, frame_shape: tuple) -> np.ndarray:
|
||||
if mask.shape[0] == 0:
|
||||
logging.warning("Empty mask received")
|
||||
return np.zeros((frame_shape[0], frame_shape[1]), dtype="uint8")
|
||||
|
||||
colors = [
|
||||
[255, 0, 255], # Purple
|
||||
[0, 255, 255], # Yellow
|
||||
[255, 255, 0], # Cyan
|
||||
[0, 255, 0], # Green
|
||||
[255, 0, 0], # Blue
|
||||
]
|
||||
|
||||
combined_colored_mask = np.zeros((frame_shape[0], frame_shape[1], 4), dtype="uint8")
|
||||
|
||||
for i in range(mask.shape[0]):
|
||||
current_mask = (mask[i, 0] > 0).cpu().numpy().astype("uint8") * 255
|
||||
if current_mask.shape[:2] != frame_shape[:2]:
|
||||
current_mask = cv2.resize(current_mask, (frame_shape[1], frame_shape[0]))
|
||||
|
||||
# Create BGRA mask with transparency
|
||||
colored_mask = np.zeros((frame_shape[0], frame_shape[1], 4), dtype="uint8")
|
||||
color = colors[i % len(colors)]
|
||||
colored_mask[current_mask > 0] = color + [128] # Add alpha value of 128
|
||||
|
||||
# Alpha blend with existing masks
|
||||
alpha = colored_mask[:, :, 3:4] / 255.0
|
||||
combined_colored_mask = (1 - alpha) * combined_colored_mask + alpha * colored_mask
|
||||
|
||||
# Convert back to BGR for display
|
||||
combined_colored_mask = combined_colored_mask[:, :, :3].astype("uint8")
|
||||
return combined_colored_mask
|
||||
|
||||
def segment_images(
|
||||
self,
|
||||
images,
|
||||
@@ -171,6 +209,8 @@ class Sam2RealtimeSegmentation:
|
||||
# bboxes=None,
|
||||
# individual_objects=False,
|
||||
# mask=None,
|
||||
threshold=0.5,
|
||||
show_point=False,
|
||||
):
|
||||
model = sam2_model["model"]
|
||||
device = sam2_model["device"]
|
||||
@@ -180,21 +220,17 @@ class Sam2RealtimeSegmentation:
|
||||
|
||||
processed_frames = []
|
||||
mask_list = []
|
||||
# The `model` variable is now ready and equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor
|
||||
# The `model` is equivalent to `predictor` returned by sam2.build_sam.build_sam2_camera_predictor
|
||||
if self.predictor is None:
|
||||
self.predictor = model
|
||||
|
||||
def process_frame(frame, frame_idx):
|
||||
with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16):
|
||||
frame = frame.to(device).float() # Keep everything in torch
|
||||
frame = frame.to(device).float()
|
||||
|
||||
if not self.if_init:
|
||||
self.predictor.load_first_frame(frame)
|
||||
self.if_init = True
|
||||
# obj_id = 1
|
||||
# point = [256, 256]
|
||||
# points = [point]
|
||||
# labels = [1]
|
||||
|
||||
coordinates_positive_list = ast.literal_eval(coordinates_positive)
|
||||
point_labels_list = ast.literal_eval(point_labels)
|
||||
@@ -208,34 +244,34 @@ class Sam2RealtimeSegmentation:
|
||||
points=[point_tuple],
|
||||
labels=[point_labels_list[idx]]
|
||||
)
|
||||
|
||||
# _, _, _ = self.predictor.add_new_prompt(frame_idx, obj_id, points=points, labels=labels)
|
||||
else:
|
||||
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 = (out_mask_logits[0, 0] > threshold).byte()
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask.unsqueeze(0).unsqueeze(0).float(),
|
||||
size=(frame.shape[0], frame.shape[1]),
|
||||
mode='nearest'
|
||||
).squeeze(0).squeeze(0).byte().to(device) # Move the interpolated mask to the correct device
|
||||
).squeeze(0).squeeze(0).byte() # Move the interpolated mask to the correct device
|
||||
else:
|
||||
mask = torch.ones((frame.shape[0], frame.shape[1]), device=device, dtype=torch.uint8)
|
||||
|
||||
# Ensure frame is on the same device
|
||||
frame = frame.to(device)
|
||||
automask_colored = self._process_mask(mask,frame.shape)
|
||||
|
||||
mask_colored = torch.stack([mask] * 3, dim=2).to(device) # Create 3-channel mask and move to device
|
||||
# Draw points on the mask
|
||||
if show_point:
|
||||
for point in coordinates_positive:
|
||||
cv2.circle(automask_colored, tuple(point), radius=5, color=(0, 0, 255), thickness=-1)
|
||||
|
||||
overlayed_frame = torch.add(frame * 0.7, mask_colored * 0.3).to(device)
|
||||
processed_frames.append(overlayed_frame)
|
||||
automasked_frame = torch.add(frame * 0.7, automask_colored * 0.3)
|
||||
processed_frames.append(automasked_frame)
|
||||
|
||||
constructed_mask = torch.add(frame * 0.1, mask_colored * 0.9).to(device)
|
||||
# TODO: This "mask" should be 1 channel to be returned as MASK type
|
||||
constructed_mask = torch.add(frame * 0.1, mask * 0.9)
|
||||
mask_list.append(constructed_mask)
|
||||
|
||||
|
||||
for frame_idx, img in enumerate(images):
|
||||
process_frame(img, frame_idx)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user