Add Sam3GetObjectIds and Sam3GetObjectMask for video seg
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
- 首次发布
|
||||
- 支持文本提示的图像分割
|
||||
- 视频跟踪和分割
|
||||
|
||||
@@ -10,6 +10,8 @@ class Sam3Extension(ComfyExtension):
|
||||
Sam3VideoSegmentation,
|
||||
Sam3VideoModelExtraConfig,
|
||||
Sam3Visualization,
|
||||
Sam3GetObjectIds,
|
||||
Sam3GetObjectMask,
|
||||
StringToBBox
|
||||
]
|
||||
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user