diff --git a/README.md b/README.md index d792850..b6339c7 100644 --- a/README.md +++ b/README.md @@ -101,7 +101,7 @@ Get all object IDs and count from SAM3 Video Segmentation output. - `objects`: Objects output from SAM3 Video Segmentation node (contains `obj_ids` and `obj_masks`) **Outputs:** -- `obj_ids`: List of all object IDs that were tracked in the video +- `object_ids`: List of all object IDs that were tracked in the video - `count`: Total number of objects tracked **Use Case:** @@ -123,7 +123,7 @@ Extract mask for a specific object ID from SAM3 Video Segmentation output. **Inputs:** - `objects`: Objects output from SAM3 Video Segmentation node (contains `obj_ids` and `obj_masks`) -- `obj_id`: Object ID to extract mask for (min: 0, max: 1000, default: 1) +- `obj_id`: Object Index to extract mask for (min: 0, max: 1000, default: 0) **Outputs:** - `mask`: Extracted mask tensor for the specified object ID [B, H, W] format. Returns empty mask if object ID not found diff --git a/README_CN.md b/README_CN.md index 1e98284..3c46bc2 100644 --- a/README_CN.md +++ b/README_CN.md @@ -101,7 +101,7 @@ - `objects`:来自 SAM3 视频分割节点的对象输出(包含 `obj_ids` 和 `obj_masks`) **输出:** -- `obj_ids`:视频中跟踪到的所有对象 ID 列表 +- `object_ids`:视频中跟踪到的所有对象 ID 列表 - `count`:跟踪到的对象总数 **使用场景:** @@ -123,7 +123,7 @@ **输入:** - `objects`:来自 SAM3 视频分割节点的对象输出(包含 `obj_ids` 和 `obj_masks`) -- `obj_id`:要提取遮罩的对象 ID(最小值:0,最大值:1000,默认:1) +- `obj_id`:要提取遮罩的对象索引ID(最小值:0,最大值:1000,默认:0) **输出:** - `mask`:指定对象 ID 的提取遮罩张量,[B, H, W] 格式。如果未找到对象 ID,则返回空遮罩 diff --git a/locales/en/nodeDefs.json b/locales/en/nodeDefs.json index 34e2a45..86e2523 100644 --- a/locales/en/nodeDefs.json +++ b/locales/en/nodeDefs.json @@ -317,7 +317,7 @@ }, "obj_id": { "name": "obj_id", - "tooltip": "Object ID to extract mask for (0-1000)" + "tooltip": "Object Index to extract mask for (0-1000)" } }, "outputs": { diff --git a/locales/zh/nodeDefs.json b/locales/zh/nodeDefs.json index 8c97d56..0604124 100644 --- a/locales/zh/nodeDefs.json +++ b/locales/zh/nodeDefs.json @@ -324,14 +324,14 @@ "tooltip": "来自 SAM3 视频分割节点的对象输出" }, "obj_id": { - "name": "对象 ID", - "tooltip": "要提取遮罩的对象 ID(0-1000)" + "name": "对象索引ID", + "tooltip": "要提取遮罩的对象索引ID(0-1000)" } }, "outputs": { "0": { "name": "遮罩", - "tooltip": "指定对象 ID 的提取遮罩 [B, H, W] 格式" + "tooltip": "指定对象索引ID 的提取遮罩 [B, H, W] 格式" } } }, diff --git a/nodes.py b/nodes.py index d5af0bc..79d61ba 100644 --- a/nodes.py +++ b/nodes.py @@ -721,8 +721,6 @@ class Sam3VideoSegmentation(io.ComfyNode): } # Use dictionary to store object_masks by frame_idx to handle non-sequential frame processing object_masks_dict = {} - # Store obj_masks for each frame in a dictionary - obj_masks_by_frame = {} for response in video_predictor.handle_stream_request( request=dict( @@ -743,18 +741,16 @@ class Sam3VideoSegmentation(io.ComfyNode): mask = outputs["out_binary_masks"] # Store mask for this frame if mask.shape[0] > 0: - obj_masks_by_frame[frame_idx] = mask - # Convert mask to tensor and store by frame_idx - mask_tensor = torch.from_numpy(mask).float() - object_masks_dict[frame_idx] = mask_tensor + # Store numpy array in object_masks_dict for consistent processing + object_masks_dict[frame_idx] = mask merged_mask = np.any(mask, axis=0).astype(np.float32) frame_masks = torch.from_numpy(merged_mask) output_masks[frame_idx] = frame_masks else: - object_masks_dict[frame_idx] = torch.zeros((1, H, W)) + object_masks_dict[frame_idx] = np.zeros((1, H, W), dtype=np.float32) else: - object_masks_dict[frame_idx] = torch.zeros((1, H, W)) + object_masks_dict[frame_idx] = np.zeros((1, H, W), dtype=np.float32) # Update progress bar processed_frames += 1 @@ -778,53 +774,40 @@ class Sam3VideoSegmentation(io.ComfyNode): if not keep_model_loaded and close_after_propagation: video_predictor.shutdown() - # Convert obj_masks_by_frame to ordered list matching frame indices - if len(obj_masks_by_frame) > 0: - # Find maximum number of objects across all processed frames - max_num_objects = max(mask.shape[0] for mask in obj_masks_by_frame.values()) - - # Create ordered list of obj_masks by frame index - ordered_obj_masks = [] - for frame_idx in range(B): - if frame_idx in obj_masks_by_frame: - mask = obj_masks_by_frame[frame_idx] - # Pad if needed to match max_num_objects - if mask.shape[0] < max_num_objects: - padding = np.zeros((max_num_objects - mask.shape[0], H, W), dtype=np.float32) - mask = np.concatenate([mask, padding], axis=0) - ordered_obj_masks.append(mask) - else: - # Frame not processed, add empty mask array with correct shape - ordered_obj_masks.append(np.zeros((max_num_objects, H, W), dtype=np.float32)) - object_outputs["obj_masks"] = ordered_obj_masks - # Convert object_masks_dict to ordered list and pad to have same number of objects across all frames if len(object_masks_dict) > 0: # Find the maximum number of objects across all frames max_num_objects = max(mask.shape[0] for mask in object_masks_dict.values()) # Create ordered list of masks by frame index, ensuring all B frames are included + ordered_obj_masks = [] padded_masks = [] for frame_idx in range(B): if frame_idx in object_masks_dict: - mask = object_masks_dict[frame_idx] + mask = object_masks_dict[frame_idx] # numpy array num_objects = mask.shape[0] if num_objects < max_num_objects: - # Pad with zero masks - padding = torch.zeros((max_num_objects - num_objects, H, W)) - padded_mask = torch.cat([mask, padding], dim=0) - padded_masks.append(padded_mask) + # Pad with zero masks (numpy for obj_masks) + padding = np.zeros((max_num_objects - num_objects, H, W), dtype=np.float32) + padded_mask = np.concatenate([mask, padding], axis=0) + ordered_obj_masks.append(padded_mask) + padded_masks.append(torch.from_numpy(padded_mask)) else: - padded_masks.append(mask) + ordered_obj_masks.append(mask) + padded_masks.append(torch.from_numpy(mask)) else: # Frame not processed, add empty mask with correct shape + empty_mask = np.zeros((max_num_objects, H, W), dtype=np.float32) + ordered_obj_masks.append(empty_mask) padded_masks.append(torch.zeros((max_num_objects, H, W))) # Now stack all B frames object_masks = torch.stack(padded_masks, dim=0) + object_outputs["obj_masks"] = ordered_obj_masks else: # No masks detected, create empty tensor object_masks = torch.zeros((B, 1, H, W)) + object_outputs["obj_masks"] = [] return io.NodeOutput(output_masks, session_id, object_outputs, object_masks) @@ -1146,8 +1129,8 @@ class Sam3GetObjectIds(io.ComfyNode): ], outputs=[ io.Int.Output( - "obj_ids", - display_name="obj_ids", + "object_ids", + display_name="object_ids", tooltip="Comma-separated list of all object IDs" ), io.Int.Output( @@ -1169,7 +1152,7 @@ class Sam3GetObjectIds(io.ComfyNode): - 'obj_masks': list of numpy arrays for each frame Returns: - obj_ids_str: Comma-separated string of all object IDs + object_ids: all object IDs count: Total number of objects """ if objects is None: @@ -1194,7 +1177,7 @@ class Sam3GetObjectIds(io.ComfyNode): class Sam3GetObjectMask(io.ComfyNode): - """Extract mask for a specific object ID from Sam3VideoSegmentation output.""" + """Extract mask for a specific object index from Sam3VideoSegmentation output.""" @classmethod def define_schema(cls): @@ -1202,7 +1185,7 @@ class Sam3GetObjectMask(io.ComfyNode): node_id="easy sam3GetObjectMask", display_name="SAM3 Get Object Mask", category="EasyUse/Sam3", - description="Extract mask for a specific object ID from Sam3VideoSegmentation objects output", + description="Extract mask for a specific object index from Sam3VideoSegmentation objects output", inputs=[ io.Custom(io_type="EASY_SAM3_OBJECTS_OUTPUT").Input( "objects", @@ -1211,78 +1194,84 @@ class Sam3GetObjectMask(io.ComfyNode): ), io.Int.Input( "obj_id", - default=1, + default=0, min=0, max=1000, - tooltip="Object ID to extract mask for" + tooltip="Object index (0-based) to extract mask for, not the actual object ID" ), ], outputs=[ io.Mask.Output( "mask", display_name="mask", - tooltip="Extracted mask for the specified object ID" - ), + tooltip="Extracted mask for the specified object index" + ) ] ) @classmethod def execute(cls, objects, obj_id) -> io.NodeOutput: """ - Extract mask for a specific object ID from objects output. + Extract mask for a specific object index from objects output. Args: objects: Dictionary containing: - 'obj_ids': numpy array of object IDs [num_objects] - 'obj_masks': list of numpy arrays, each [num_objects, H, W] for each frame - obj_id: Object ID to extract mask for + obj_id: Object index (0-based) to extract mask for Returns: - Batch of masks tensor [num_frames, H, W] for the specified object ID + mask: Batch of masks tensor [num_frames, H, W] for the specified object index """ if objects is None: raise ValueError("Objects input cannot be None") - obj_ids = objects.get("obj_ids", None) obj_masks = objects.get("obj_masks", None) + obj_ids = objects.get("obj_ids", None) - if obj_ids is None or obj_masks is None: - raise ValueError("Objects must contain both 'obj_ids' and 'obj_masks' keys") + if obj_masks is None: + raise ValueError("Objects must contain 'obj_masks' key") + + if obj_ids is None: + raise ValueError("Objects must contain 'obj_ids' key") - # Convert obj_ids to numpy array if needed - if isinstance(obj_ids, torch.Tensor): - obj_ids = obj_ids.cpu().numpy() - - # Find the index of the requested obj_id + # Use obj_idx directly as the index try: - obj_index = np.nonzero(obj_ids == obj_id)[0] - - if len(obj_index) == 0: - logger.warning(f"Object ID {obj_id} not found in objects. Available IDs: {obj_ids}") - # Return empty masks for all frames - if isinstance(obj_masks, list) and len(obj_masks) > 0: - first_frame = obj_masks[0] - if isinstance(first_frame, np.ndarray) and len(first_frame.shape) >= 2: - H, W = first_frame.shape[-2], first_frame.shape[-1] - num_frames = len(obj_masks) - empty_masks = torch.zeros((num_frames, H, W), dtype=torch.float32) - else: - empty_masks = torch.zeros((1, 1, 1), dtype=torch.float32) - else: - empty_masks = torch.zeros((1, 1, 1), dtype=torch.float32) + if not isinstance(obj_masks, list) or len(obj_masks) == 0: + logger.warning("obj_masks is empty or invalid") + empty_masks = torch.zeros((1, 1, 1), dtype=torch.float32) return io.NodeOutput(empty_masks) - obj_index = obj_index[0] + # Get the first frame to check dimensions + first_frame = obj_masks[0] + if isinstance(first_frame, torch.Tensor): + first_frame = first_frame.cpu().numpy() + + num_objects = first_frame.shape[0] if len(first_frame.shape) >= 3 else 0 + + # Validate obj_idx index + if obj_id < 0 or obj_id >= num_objects: + logger.warning(f"Object index {obj_id} out of range. Available indices: 0-{num_objects-1}") + # Return empty masks for all frames + H, W = first_frame.shape[-2], first_frame.shape[-1] + num_frames = len(obj_masks) + empty_masks = torch.zeros((num_frames, H, W), dtype=torch.float32) + return io.NodeOutput(empty_masks, -1) + + # Get the actual object ID for this index + if isinstance(obj_ids, torch.Tensor): + obj_ids = obj_ids.cpu().numpy() + object_id = int(obj_ids[obj_id]) - # Extract masks for this object across all frames + # Extract masks for this object index across all frames extracted_masks = [] for frame_masks in obj_masks: # frame_masks is [num_objects, H, W] if isinstance(frame_masks, torch.Tensor): frame_masks = frame_masks.cpu().numpy() - # Extract the mask for this object in this frame - obj_mask = frame_masks[obj_index] + # Extract the mask for this object index in this frame + obj_mask = frame_masks[obj_id] # Convert boolean mask to float if obj_mask.dtype == bool: @@ -1294,13 +1283,13 @@ class Sam3GetObjectMask(io.ComfyNode): masks_array = np.stack(extracted_masks, axis=0) mask_tensor = torch.from_numpy(masks_array).float() - logger.info(f"Extracted masks for object ID {obj_id} with shape {mask_tensor.shape} ({len(obj_masks)} frames)") + logger.info(f"Extracted masks for object index {obj_id} (ID: {object_id}) with shape {mask_tensor.shape} ({len(obj_masks)} frames)") return io.NodeOutput(mask_tensor) except Exception as e: logger.error(f"Error extracting object mask: {str(e)}") - raise ValueError(f"Error extracting object mask for ID {obj_id}: {str(e)}") + raise ValueError(f"Error extracting object mask for index {obj_id}: {str(e)}") class StringToBBox(io.ComfyNode): diff --git a/pyproject.toml b/pyproject.toml index 67bd9b4..b3d91d4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-easy-sam3" description = "A ComfyUI custom node package for segment anything 3" -version = "1.0.2" +version = "1.0.3" license = {file = "LICENSE"} dependencies = ["timm>=1.0.17", "ftfy==6.1.1", "regex", "iopath>=0.1.10", "einops>=0.6.0", "decord>=0.6.0", "pycocotools>=2.0.10", "numpy>=1.26", "tqdm", "typing_extensions"]