wip fixup automasking, show coord point

This commit is contained in:
Peter Schroedl
2024-11-30 18:18:52 +01:00
parent a5ad3b699f
commit cc4caff286
+55 -19
View File
@@ -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)