Add Sam3GetObjectIds and Sam3GetObjectMask for video seg

This commit is contained in:
yolain
2025-11-26 02:29:21 +08:00
parent ba6336ac94
commit 129d139191
7 changed files with 371 additions and 7 deletions
+50 -2
View File
@@ -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
+50 -2
View File
@@ -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
- 首次发布
- 支持文本提示的图像分割
- 视频跟踪和分割
+2
View File
@@ -10,6 +10,8 @@ class Sam3Extension(ComfyExtension):
Sam3VideoSegmentation,
Sam3VideoModelExtraConfig,
Sam3Visualization,
Sam3GetObjectIds,
Sam3GetObjectMask,
StringToBBox
]
+38
View File
@@ -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": {
+38
View File
@@ -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": {
+192 -2
View File
@@ -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."""
+1 -1
View File
@@ -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"]