Fix object indexes failing to match object IDs

This commit is contained in:
yolain
2025-12-08 22:36:53 +08:00
parent 851a5be3c0
commit b0c9f6b131
6 changed files with 73 additions and 84 deletions
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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,则返回空遮罩
+1 -1
View File
@@ -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": {
+3 -3
View File
@@ -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] 格式"
}
}
},
+64 -75
View File
@@ -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):
+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.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"]