From 129d139191780dcbd0c8526c6a4e32cdb02e47fa Mon Sep 17 00:00:00 2001 From: yolain Date: Wed, 26 Nov 2025 02:29:21 +0800 Subject: [PATCH] Add `Sam3GetObjectIds` and `Sam3GetObjectMask` for video seg --- README.md | 52 ++++++++++- README_CN.md | 52 ++++++++++- __init__.py | 2 + locales/en/nodeDefs.json | 38 ++++++++ locales/zh/nodeDefs.json | 38 ++++++++ nodes.py | 194 ++++++++++++++++++++++++++++++++++++++- pyproject.toml | 2 +- 7 files changed, 371 insertions(+), 7 deletions(-) diff --git a/README.md b/README.md index ce16570..526bba9 100644 --- a/README.md +++ b/README.md @@ -91,8 +91,49 @@ Track and segment objects across video frames with advanced prompting options. - `masks`: Tracked segmentation masks for all frames [B, H, W] format with merged object masks per frame - `session_id`: Session ID string for resuming tracking in subsequent calls - `objects`: Object tracking information dictionary containing `obj_ids` and `obj_masks` arrays +- `obj_masks`: Individual object masks per frame [B, N, H, W] format where N is the maximum number of objects tracked -### 4. SAM3 Video Model Extra Config +### 4. SAM3 Get Object IDs +Get all object IDs and count from SAM3 Video Segmentation output. + +**Inputs:** +- `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 +- `count`: Total number of objects tracked + +**Use Case:** +This node is essential for understanding what objects were detected and tracked in your video. It helps you: +- Discover all object IDs available for extraction +- Determine how many objects were successfully tracked +- Plan downstream processing based on the number of tracked objects +- Debug tracking issues by verifying which objects were detected + +**Example:** +If your video tracked 3 objects, this node will output: +- `obj_ids`: [0, 1, 2] +- `count`: 3 + +You can then use these IDs with the SAM3 Get Object Mask node to extract individual object masks. + +### 6. SAM3 Get Object Mask +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) + +**Outputs:** +- `mask`: Extracted mask tensor for the specified object ID [B, H, W] format. Returns empty mask if object ID not found + +**Use Case:** +This node is useful when you have multiple tracked objects in a video and want to isolate a specific object's mask for further processing. For example: +- Extract a person's mask (object_id=1) from a video with multiple people +- Process different objects separately in your workflow +- Apply different effects or transformations to specific tracked objects + +### 7. SAM3 Video Model Extra Config Configure advanced parameters for video segmentation to fine-tune tracking behavior. **Parameters:** @@ -117,7 +158,7 @@ Configure advanced parameters for video segmentation to fine-tune tracking behav **Output:** - `extra_config`: Configuration dictionary for Video Segmentation node -### 5. Sam3 Visualization +### 8. Sam3 Visualization Visualize segmentation masks with bounding boxes and confidence scores overlaid on images. **Inputs:** @@ -180,7 +221,14 @@ Contributions are welcome! Please feel free to submit issues or pull requests. ## Changelog +### v1.0.1 + +- Added `easy sam3GetObjectIds` node to get object ID list +- Added `easy sam3GetObjectMask` node to extract mask for specific object ID +- Fixed object masks not aligned with frames when `start_frame_index` is not 0 + ### v1.0.0 + - Initial release - Image segmentation with text prompts - Video tracking and segmentation diff --git a/README_CN.md b/README_CN.md index bf73c46..b7ed972 100644 --- a/README_CN.md +++ b/README_CN.md @@ -91,8 +91,49 @@ - `masks`:所有帧的跟踪分割遮罩,[B, H, W] 格式,每帧的对象遮罩已合并 - `session_id`:会话 ID 字符串,用于后续调用中恢复跟踪 - `objects`:对象跟踪信息字典,包含 `obj_ids` 和 `obj_masks` 数组 +- `obj_masks`:每帧的单个对象遮罩,[B, N, H, W] 格式,其中 N 是跟踪的最大对象数量 -### 4. SAM3 视频模型额外配置 +### 4. SAM3 获取对象ID列表 +从 SAM3 视频分割输出中获取所有对象 ID 和数量。 + +**输入:** +- `objects`:来自 SAM3 视频分割节点的对象输出(包含 `obj_ids` 和 `obj_masks`) + +**输出:** +- `obj_ids`:视频中跟踪到的所有对象 ID 列表 +- `count`:跟踪到的对象总数 + +**使用场景:** +此节点对于了解视频中检测和跟踪到的对象至关重要。它可以帮助您: +- 发现可用于提取的所有对象 ID +- 确定成功跟踪的对象数量 +- 根据跟踪对象的数量规划下游处理 +- 通过验证检测到的对象来调试跟踪问题 + +**示例:** +如果您的视频跟踪了 3 个对象,此节点将输出: +- `obj_ids`: [0, 1, 2] +- `count`: 3 + +然后您可以使用这些 ID 配合 SAM3 获取对象遮罩节点来提取单个对象的遮罩。 + +### 6. SAM3 获取对象遮罩 +从 SAM3 视频分割输出中提取特定对象 ID 的遮罩。 + +**输入:** +- `objects`:来自 SAM3 视频分割节点的对象输出(包含 `obj_ids` 和 `obj_masks`) +- `obj_id`:要提取遮罩的对象 ID(最小值:0,最大值:1000,默认:1) + +**输出:** +- `mask`:指定对象 ID 的提取遮罩张量,[B, H, W] 格式。如果未找到对象 ID,则返回空遮罩 + +**使用场景:** +当视频中有多个跟踪对象时,此节点可用于隔离特定对象的遮罩以进行进一步处理。例如: +- 从包含多人的视频中提取一个人的遮罩(object_id=1) +- 在工作流中单独处理不同的对象 +- 对特定跟踪对象应用不同的效果或变换 + +### 7. SAM3 视频模型额外配置 配置视频分割的高级参数,以微调跟踪行为。 **参数:** @@ -117,7 +158,7 @@ **输出:** - `extra_config`:视频分割节点的配置字典 -### 5. Sam3 可视化生成 +### 8. Sam3 可视化生成 在图像上可视化分割遮罩,并叠加边界框和置信度分数。 **输入:** @@ -180,7 +221,14 @@ ## 更新日志 +### v1.0.1 + +- 增加 `easy sam3GetObjectIds` 节点以获取对象ID列表 +- 增加 `easy sam3GetObjectMask` 节点以提取特定对象ID的遮罩 +- 修复 `start_frame_index` 不为0时,对象遮罩未对齐帧 + ### v1.0.0 + - 首次发布 - 支持文本提示的图像分割 - 视频跟踪和分割 diff --git a/__init__.py b/__init__.py index 78524ef..b0b96dc 100644 --- a/__init__.py +++ b/__init__.py @@ -10,6 +10,8 @@ class Sam3Extension(ComfyExtension): Sam3VideoSegmentation, Sam3VideoModelExtraConfig, Sam3Visualization, + Sam3GetObjectIds, + Sam3GetObjectMask, StringToBBox ] diff --git a/locales/en/nodeDefs.json b/locales/en/nodeDefs.json index df5ef05..93ab0ad 100644 --- a/locales/en/nodeDefs.json +++ b/locales/en/nodeDefs.json @@ -289,6 +289,44 @@ } } }, + "easy sam3GetObjectIds": { + "display_name": "SAM3 Get Object IDs", + "inputs": { + "objects": { + "name": "objects", + "tooltip": "Objects output from Sam3VideoSegmentation node" + } + }, + "outputs": { + "0": { + "name": "obj_ids", + "tooltip": "Comma-separated list of all object IDs" + }, + "1": { + "name": "count", + "tooltip": "Total number of objects" + } + } + }, + "easy sam3GetObjectMask": { + "display_name": "SAM3 Get Object Mask", + "inputs": { + "objects": { + "name": "objects", + "tooltip": "Objects output from Sam3VideoSegmentation node" + }, + "obj_id": { + "name": "obj_id", + "tooltip": "Object ID to extract mask for (0-1000)" + } + }, + "outputs": { + "0": { + "name": "mask", + "tooltip": "Extracted mask for the specified object ID [B, H, W] format" + } + } + }, "easy stringToBBox": { "display_name": "String to BBox", "inputs": { diff --git a/locales/zh/nodeDefs.json b/locales/zh/nodeDefs.json index 3e1eb1a..e36b52d 100644 --- a/locales/zh/nodeDefs.json +++ b/locales/zh/nodeDefs.json @@ -297,6 +297,44 @@ } } }, + "easy sam3GetObjectIds": { + "display_name": "SAM3 获取对象 ID", + "inputs": { + "objects": { + "name": "对象", + "tooltip": "来自 SAM3 视频分割节点的对象输出" + } + }, + "outputs": { + "0": { + "name": "对象 ID", + "tooltip": "所有对象 ID 的逗号分隔列表" + }, + "1": { + "name": "数量", + "tooltip": "对象总数" + } + } + }, + "easy sam3GetObjectMask": { + "display_name": "SAM3 获取对象遮罩", + "inputs": { + "objects": { + "name": "对象", + "tooltip": "来自 SAM3 视频分割节点的对象输出" + }, + "obj_id": { + "name": "对象 ID", + "tooltip": "要提取遮罩的对象 ID(0-1000)" + } + }, + "outputs": { + "0": { + "name": "遮罩", + "tooltip": "指定对象 ID 的提取遮罩 [B, H, W] 格式" + } + } + }, "easy stringToBBox": { "display_name": "字符串转边界框", "inputs": { diff --git a/nodes.py b/nodes.py index 6109c57..a09840a 100644 --- a/nodes.py +++ b/nodes.py @@ -709,10 +709,12 @@ class Sam3VideoSegmentation(io.ComfyNode): object_outputs = { "obj_ids":None, - "obj_masks":None + "obj_masks":[] } # 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( @@ -731,7 +733,8 @@ class Sam3VideoSegmentation(io.ComfyNode): if outputs: if "out_binary_masks" in outputs: mask = outputs["out_binary_masks"] - object_outputs["obj_masks"] = mask + # Store mask for this frame + obj_masks_by_frame[frame_idx] = mask if mask.shape[0] > 0: # Convert mask to tensor and store by frame_idx mask_tensor = torch.from_numpy(mask).float() @@ -767,6 +770,18 @@ 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: + # 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: + ordered_obj_masks.append(obj_masks_by_frame[frame_idx]) + else: + # Frame not processed, add empty mask array + ordered_obj_masks.append(np.zeros((1, 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: # Create ordered list of masks by frame index @@ -1099,6 +1114,181 @@ class Sam3Visualization(io.ComfyNode): # Return with preview UI return io.NodeOutput(output_images,) +class Sam3GetObjectIds(io.ComfyNode): + """Get all object IDs from Sam3VideoSegmentation output.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="easy sam3GetObjectIds", + display_name="SAM3 Get Object IDs", + category="EasyUse/Sam3", + description="Get all object IDs and count from Sam3VideoSegmentation objects output", + inputs=[ + io.Custom(io_type="EASY_SAM3_OBJECTS_OUTPUT").Input( + "objects", + display_name="objects", + tooltip="Objects output from Sam3VideoSegmentation node" + ), + ], + outputs=[ + io.Int.Output( + "obj_ids", + display_name="obj_ids", + tooltip="Comma-separated list of all object IDs" + ), + io.Int.Output( + "count", + display_name="count", + tooltip="Total number of objects" + ), + ] + ) + + @classmethod + def execute(cls, objects) -> io.NodeOutput: + """ + Get all object IDs from objects output. + + Args: + objects: Dictionary containing: + - 'obj_ids': numpy array of object IDs [num_objects] + - 'obj_masks': list of numpy arrays for each frame + + Returns: + obj_ids_str: Comma-separated string of all object IDs + count: Total number of objects + """ + if objects is None: + raise ValueError("Objects input cannot be None") + + obj_ids = objects.get("obj_ids", None) + + 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() + + # Get count + count = len(obj_ids) + + # Convert to comma-separated string + obj_ids = [int(obj_id) for obj_id in obj_ids] + + return io.NodeOutput(obj_ids, count) + + +class Sam3GetObjectMask(io.ComfyNode): + """Extract mask for a specific object ID from Sam3VideoSegmentation output.""" + + @classmethod + def define_schema(cls): + return io.Schema( + 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", + inputs=[ + io.Custom(io_type="EASY_SAM3_OBJECTS_OUTPUT").Input( + "objects", + display_name="objects", + tooltip="Objects output from Sam3VideoSegmentation node" + ), + io.Int.Input( + "obj_id", + default=1, + min=0, + max=1000, + tooltip="Object ID to extract mask for" + ), + ], + outputs=[ + io.Mask.Output( + "mask", + display_name="mask", + tooltip="Extracted mask for the specified object ID" + ), + ] + ) + + @classmethod + def execute(cls, objects, obj_id) -> io.NodeOutput: + """ + Extract mask for a specific object ID 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 + + Returns: + Batch of masks tensor [num_frames, H, W] for the specified object ID + """ + 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) + + if obj_ids is None or obj_masks is None: + raise ValueError("Objects must contain both 'obj_ids' and 'obj_masks' keys") + + # 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 + 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) + return io.NodeOutput(empty_masks) + + obj_index = obj_index[0] + + # Extract masks for this object 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] + + # Convert boolean mask to float + if obj_mask.dtype == bool: + obj_mask = obj_mask.astype(np.float32) + + extracted_masks.append(obj_mask) + + # Stack all frames: [num_frames, H, W] + 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)") + + 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)}") + class StringToBBox(io.ComfyNode): """Convert string coordinates to BBOX type.""" diff --git a/pyproject.toml b/pyproject.toml index 9d1bfe6..755ad1c 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.0" +version = "1.0.1" 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"]