Fix some bug and support points segmentation

This commit is contained in:
yolain
2025-11-22 18:23:59 +08:00
parent aae63e4a5f
commit 5ffffeaf68
18 changed files with 2521 additions and 218 deletions
+43 -23
View File
@@ -19,8 +19,6 @@ This node package brings Meta's SAM3 model to ComfyUI, enabling:
![Video Segmentation Example](assets/video.png)
*Example of semantic segmentation on video frames*
> **Note**: Currently, this package supports **semantic segmentation** only.
## Features
- 🖼️ **Image Segmentation**: Segment objects using natural language prompts
@@ -45,52 +43,74 @@ Load a SAM3 model for image or video segmentation.
- `sam3_model`: Loaded SAM3 model for downstream nodes
### 2. SAM3 Image Segmentation
Segment objects in images using text prompts.
Segment objects in images using text prompts and optional geometric prompts.
**Inputs:**
- `sam3_model`: SAM3 model from Load SAM3 Model node
- `images`: Input images to segment
- `prompt`: Text description of objects to segment (e.g., "a cat", "person")
- `threshold`: Confidence threshold for detections (0.0-1.0)
- `threshold`: Confidence threshold for detections (0.0-1.0, default: 0.60)
- `keep_model_loaded`: Keep model in VRAM after inference
- `add_background`: Add background color (none, black, white, grey)
- `coordinates_positive` (optional): Positive click coordinates to refine segmentation
- `coordinates_negative` (optional): Negative click coordinates to exclude areas
- `bboxes` (optional): Bounding boxes to guide segmentation
- `mask` (optional): Input mask for refinement
**Outputs:**
- `masks`: Segmentation masks
- `images`: Segmented images (with optional background)
- `boxes`: Bounding box coordinates for detected objects
- `scores`: Confidence scores for each detection
### 3. SAM3 Video Segmentation
Track and segment objects across video frames.
Track and segment objects across video frames with advanced prompting options.
**Inputs:**
- `sam3_model`: SAM3 model in video mode
- `session_id`: Optional session ID to resume tracking
- `session_id` (optional): Session ID to resume tracking from a previous session
- `video_frames`: Video frames as image sequence
- `prompt`: Text description of objects to track
- `score_threshold_detection`: Detection confidence threshold
- `new_det_thresh`: Threshold for adding new objects
- `prompt`: Text description of objects to track (e.g., "person", "car")
- `frame_index`: Frame where initial prompt is applied (0 to max frames)
- `object_id`: Unique ID for multi-object tracking (1-1000, default: 1)
- `score_threshold_detection`: Detection confidence threshold (0.0-1.0, default: 0.5)
- `new_det_thresh`: Threshold for adding new objects (0.0-1.0, default: 0.7)
- `propagation_direction`: Propagation direction (both, forward, backward)
- `start_frame_index`: Frame index to start propagation
- `keep_model_loaded`: Keep model in VRAM
- `close_after_propagation`: Close session after completion
- `extra_config`: Additional configuration from Extra Config node
- `start_frame_index`: Frame index to start propagation (default: 0)
- `max_frames_to_track`: Maximum frames to process (-1 for all frames)
- `close_after_propagation`: Close session after completion (default: True)
- `keep_model_loaded`: Keep model in VRAM after inference
- `extra_config` (optional): Additional configuration from Extra Config node
- `positive_coords` (optional): Positive click coordinates as JSON array
- `negative_coords` (optional): Negative click coordinates as JSON array
- `bbox` (optional): Bounding box to initialize tracking
**Outputs:**
- `masks`: Tracked segmentation masks for all frames
- `session_id`: Session ID for resuming tracking
- `objects`: Object tracking information and metadata
### 4. SAM3 Video Model Extra Config
Configure advanced parameters for video segmentation.
Configure advanced parameters for video segmentation to fine-tune tracking behavior.
**Key Parameters:**
- `assoc_iou_thresh`: IoU threshold for detection-to-track matching
- `trk_assoc_iou_thresh`: Stricter IoU threshold for unmatched masklets
- `hotstart_delay`: Delay outputs to remove unmatched/duplicate tracklets
- `max_trk_keep_alive`: Maximum frames to keep track alive without detection
- `det_nms_thresh`: IoU threshold for NMS
- `fill_hole_area`: Fill holes in masks smaller than this area
- `max_num_objects`: Maximum number of objects to track
- And many more fine-tuning options...
**Parameters:**
- `assoc_iou_thresh`: IoU threshold for detection-to-track matching (0.0-1.0, default: 0.1)
- `det_nms_thresh`: IoU threshold for detection NMS (0.0-1.0, default: 0.1)
- `new_det_thresh`: Threshold for adding new objects (0.0-1.0, default: 0.7)
- `hotstart_delay`: Hold off outputs for N frames to remove unmatched/duplicate tracklets (0-100, default: 15)
- `hotstart_unmatch_thresh`: Remove tracklets unmatched for this many frames during hotstart (0-100, default: 8)
- `hotstart_dup_thresh`: Remove overlapping tracklets during hotstart (0-100, default: 8)
- `suppress_unmatched_within_hotstart`: Only suppress unmatched masks within hotstart period (default: True)
- `min_trk_keep_alive`: Minimum keep-alive value (-100-0, default: -1, negative means immediate removal)
- `max_trk_keep_alive`: Maximum frames to keep track alive without detections (0-100, default: 30)
- `init_trk_keep_alive`: Initial keep-alive when new track is created (-10-100, default: 30)
- `suppress_overlap_occlusion_thresh`: Threshold for suppressing overlapping objects (0.0-1.0, default: 0.7, 0.0 to disable)
- `suppress_det_at_boundary`: Suppress detections close to image boundaries (default: False)
- `fill_hole_area`: Fill holes in masks smaller than this area in pixels (0-1000, default: 16)
- `recondition_every_nth_frame`: Recondition tracking every N frames (-1-1000, default: 16, -1 to disable)
- `enable_masklet_confirmation`: Enable masklet confirmation to suppress unconfirmed tracklets (default: False)
- `decrease_alive_for_empty_masks`: Decrease keep-alive counter for empty masklets (default: False)
- `image_size`: Input image size for the model (256-2048, step: 8, default: 1008)
**Output:**
- `extra_config`: Configuration dictionary for Video Segmentation node
+44 -23
View File
@@ -19,7 +19,6 @@
![视频分割示例](assets/video.png)
*视频帧语义分割示例*
> **注意**:目前本插件仅支持了**语义分割**功能
## 功能特性
@@ -44,52 +43,74 @@
- `sam3_model`:已加载的 SAM3 模型,供下游节点使用
### 2. SAM3 图像分割
使用文本提示分割图像中的对象。
使用文本提示和可选的几何提示分割图像中的对象。
**输入:**
- `sam3_model`:来自"加载 SAM3 模型"节点的 SAM3 模型
- `images`:要分割的输入图像
- `prompt`:要分割的对象的文本描述(例如:"一只猫"、"人")
- `threshold`:检测的置信度阈值(0.0-1.0)
- `threshold`:检测的置信度阈值(0.0-1.0,默认:0.60)
- `keep_model_loaded`:推理后将模型保留在显存中
- `add_background`:添加背景颜色(无、黑色、白色、灰色)
- `coordinates_positive`(可选):正向点击坐标以细化分割
- `coordinates_negative`(可选):负向点击坐标以排除区域
- `bboxes`(可选):边界框来引导分割
- `mask`(可选):用于细化的输入遮罩
**输出:**
- `masks`:分割遮罩
- `images`:分割后的图像(可选背景)
- `boxes`:检测到的对象的边界框坐标
- `scores`:每个检测的置信度分数
### 3. SAM3 视频分割
跨视频帧跟踪和分割对象。
跨视频帧跟踪和分割对象,支持高级提示选项。
**输入:**
- `sam3_model`:视频模式的 SAM3 模型
- `session_id`:可选的会话 ID,用于恢复跟踪
- `session_id`(可选):会话 ID,用于从之前的会话恢复跟踪
- `video_frames`:作为图像序列的视频帧
- `prompt`:要跟踪的对象的文本描述
- `score_threshold_detection`:检测置信度阈值
- `new_det_thresh`:添加新对象的阈值
- `propagation_direction`:传播方向 (双向、前向、后向)
- `start_frame_index`:开始传播的帧索引
- `keep_model_loaded`:将模型保留在显存中
- `close_after_propagation`:完成后关闭会话
- `extra_config`:来自额外配置节点的附加配置
- `prompt`:要跟踪的对象的文本描述(例如:"人"、"汽车")
- `frame_index`:应用初始提示的帧位置(0 到最大帧数)
- `object_id`:多对象跟踪的唯一 ID(1-1000,默认:1)
- `score_threshold_detection`:检测置信度阈值(0.0-1.0,默认:0.5)
- `new_det_thresh`:添加新对象的阈值(0.0-1.0,默认:0.7)
- `propagation_direction`:传播方向(双向、前向、后向)
- `start_frame_index`:开始传播的帧索引(默认:0)
- `max_frames_to_track`:要处理的最大帧数(-1 表示所有帧)
- `close_after_propagation`:完成后关闭会话(默认:True)
- `keep_model_loaded`:推理后将模型保留在显存中
- `extra_config`(可选):来自额外配置节点的附加配置
- `positive_coords`(可选):正向点击坐标,JSON 数组格式
- `negative_coords`(可选):负向点击坐标,JSON 数组格式
- `bbox`(可选):用于初始化跟踪的边界框
**输出:**
- `masks`:所有帧的跟踪分割遮罩
- `session_id`:用于恢复跟踪的会话 ID
- `objects`:对象跟踪信息和元数据
### 4. SAM3 视频模型额外配置
配置视频分割的高级参数。
配置视频分割的高级参数,以微调跟踪行为。
**主要参数:**
- `assoc_iou_thresh`:检测到跟踪匹配的 IoU 阈值
- `trk_assoc_iou_thresh`:不匹配掩码的更严格 IoU 阈值
- `hotstart_delay`:延迟输出以移除不匹配/重复的轨迹
- `max_trk_keep_alive`:没有检测时保持跟踪活动的最大帧数
- `det_nms_thresh`:NMS 的 IoU 阈值
- `fill_hole_area`:填充遮罩中小于此面积的孔洞
- `max_num_objects`:要跟踪的最大对象数
- 还有更多微调选项...
**参数:**
- `assoc_iou_thresh`:检测到跟踪匹配的 IoU 阈值(0.0-1.0,默认:0.1)
- `det_nms_thresh`:检测 NMS 的 IoU 阈值(0.0-1.0,默认:0.1)
- `new_det_thresh`:添加新对象的阈值(0.0-1.0,默认:0.7)
- `hotstart_delay`:延迟 N 帧输出以移除不匹配/重复的轨迹(0-100,默认:15)
- `hotstart_unmatch_thresh`:在热启动期间移除未匹配此帧数的轨迹(0-100,默认:8)
- `hotstart_dup_thresh`:在热启动期间移除重叠的轨迹(0-100,默认:8)
- `suppress_unmatched_within_hotstart`:仅在热启动期间抑制未匹配的遮罩(默认:True)
- `min_trk_keep_alive`:最小保持活动值(-100-0,默认:-1,负值表示立即移除)
- `max_trk_keep_alive`:没有检测时保持跟踪活动的最大帧数(0-100,默认:30)
- `init_trk_keep_alive`:创建新跟踪时的初始保持活动值(-10-100,默认:30)
- `suppress_overlap_occlusion_thresh`:基于近期遮挡抑制重叠对象的阈值(0.0-1.0,默认:0.7,0.0 表示禁用)
- `suppress_det_at_boundary`:抑制接近图像边界的检测(默认:False)
- `fill_hole_area`:填充遮罩中小于此面积(像素)的孔洞(0-1000,默认:16)
- `recondition_every_nth_frame`:每 N 帧重新调整跟踪(-1-1000,默认:16,-1 表示禁用)
- `enable_masklet_confirmation`:启用掩码确认以抑制未确认的轨迹(默认:False)
- `decrease_alive_for_empty_masks`:减少空掩码的保持活动计数器(默认:False)
- `image_size`:模型的输入图像大小(256-2048,步长:8,默认:1008)
**输出:**
- `extra_config`:视频分割节点的配置字典
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 259 KiB

After

Width:  |  Height:  |  Size: 784 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 251 KiB

After

Width:  |  Height:  |  Size: 688 KiB

File diff suppressed because one or more lines are too long
@@ -0,0 +1,281 @@
{
"id": "307743a2-8a4b-4090-b8dc-3125ff31fa14",
"revision": 0,
"last_node_id": 223,
"last_link_id": 304,
"nodes": [
{
"id": 222,
"type": "PreviewImage",
"pos": [
-3883.9588057841556,
593.0894010290336
],
"size": [
339.2921196124935,
246
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 303
}
],
"outputs": [],
"properties": {
"cnr_id": "comfy-core",
"ver": "0.3.70",
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 218,
"type": "MaskPreview+",
"pos": [
-3880.8437110898844,
280.9015183807766
],
"size": [
332.95740427927046,
258
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "mask",
"type": "MASK",
"link": 298
}
],
"outputs": [],
"properties": {
"cnr_id": "comfyui_essentials",
"ver": "9d9f4bedfc9f0321c19faf71855e228c93bd0dc9",
"Node name for S&R": "MaskPreview+"
},
"widgets_values": []
},
{
"id": 190,
"type": "easy sam3ModelLoader",
"pos": [
-4235.041825672723,
334.7519653807837
],
"size": [
314.09852906952574,
130
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "sam3_model",
"type": "EASY_SAM3_MODEL",
"links": [
282
]
}
],
"properties": {
"Node name for S&R": "easy sam3ModelLoader"
},
"widgets_values": [
"sam3.pt",
"image",
"cuda",
"fp16"
]
},
{
"id": 221,
"type": "LoadImage",
"pos": [
-4678.319118164902,
372.4935004943305
],
"size": [
406.60579628243977,
446.341753250938
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
302
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"cnr_id": "comfy-core",
"ver": "0.3.70",
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"1 (3).png",
"image"
]
},
{
"id": 213,
"type": "easy sam3ImageSegmentation",
"pos": [
-4241.387303264146,
538.9760914950197
],
"size": [
337.2438563422779,
281.966544722883
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "sam3_model",
"type": "EASY_SAM3_MODEL",
"link": 282
},
{
"name": "images",
"type": "IMAGE",
"link": 302
},
{
"name": "coordinates_positive",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "coordinates_negative",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "bboxes",
"shape": 7,
"type": "BBOX",
"link": null
},
{
"name": "mask",
"shape": 7,
"type": "MASK",
"link": null
}
],
"outputs": [
{
"name": "masks",
"shape": 6,
"type": "MASK",
"links": [
298
]
},
{
"name": "images",
"shape": 6,
"type": "IMAGE",
"links": [
303
]
},
{
"name": "boxes",
"shape": 6,
"type": "STRING",
"links": null
},
{
"name": "scores",
"shape": 6,
"type": "STRING",
"links": null
}
],
"properties": {
"Node name for S&R": "easy sam3ImageSegmentation"
},
"widgets_values": [
"apple",
0.4,
false,
"none"
]
}
],
"links": [
[
282,
190,
0,
213,
0,
"EASY_SAM3_MODEL"
],
[
298,
213,
0,
218,
0,
"MASK"
],
[
302,
221,
0,
213,
1,
"IMAGE"
],
[
303,
213,
1,
222,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 1.3920191078517683,
"offset": [
4898.859436279234,
-79.3314399803767
]
},
"workflowRendererVersion": "LG",
"frontendVersion": "1.33.5",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
File diff suppressed because one or more lines are too long
@@ -0,0 +1,382 @@
{
"id": "307743a2-8a4b-4090-b8dc-3125ff31fa14",
"revision": 0,
"last_node_id": 232,
"last_link_id": 319,
"nodes": [
{
"id": 228,
"type": "easy sam3ModelLoader",
"pos": [
-4519.153385861623,
232.95641658512474
],
"size": [
270,
130
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "sam3_model",
"type": "EASY_SAM3_MODEL",
"links": [
312
]
}
],
"properties": {
"Node name for S&R": "easy sam3ModelLoader"
},
"widgets_values": [
"sam3.pt",
"video",
"cuda",
"fp16"
]
},
{
"id": 231,
"type": "MaskToImage",
"pos": [
-4446.76917913515,
119.15114245007179
],
"size": [
184.6236328125,
26
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "mask",
"type": "MASK",
"link": 318
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
319
]
}
],
"properties": {
"cnr_id": "comfy-core",
"ver": "0.3.70",
"Node name for S&R": "MaskToImage"
}
},
{
"id": 232,
"type": "VHS_VideoCombine",
"pos": [
-4186.105715232003,
-114.75297305036214
],
"size": [
422.65873186131785,
1045.941801892951
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 319
},
{
"name": "audio",
"shape": 7,
"type": "AUDIO",
"link": null
},
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"cnr_id": "comfyui-videohelpersuite",
"ver": "8923bd836bdab8b7bbdf4ed104b7d045e70c66e2",
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "AnimateDiff",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": false,
"trim_to_audio": false,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "AnimateDiff_00003.mp4",
"subfolder": "",
"type": "temp",
"format": "video/h264-mp4",
"frame_rate": 24,
"workflow": "AnimateDiff_00003.png",
"fullpath": "E:\\ComfyUI_windows_portable\\ComfyUI\\temp\\AnimateDiff_00003.mp4"
}
}
}
},
{
"id": 226,
"type": "VHS_LoadVideo",
"pos": [
-4808.758351453994,
184.48164872396106
],
"size": [
237.6240234375,
701.2149739583333
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
313
]
},
{
"name": "frame_count",
"type": "INT",
"links": null
},
{
"name": "audio",
"type": "AUDIO",
"links": null
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null
}
],
"properties": {
"cnr_id": "comfyui-videohelpersuite",
"ver": "8923bd836bdab8b7bbdf4ed104b7d045e70c66e2",
"Node name for S&R": "VHS_LoadVideo"
},
"widgets_values": {
"video": "时尚走秀.mp4",
"force_rate": 24,
"custom_width": 480,
"custom_height": 832,
"frame_load_cap": 121,
"skip_first_frames": 0,
"select_every_nth": 1,
"format": "AnimateDiff",
"choose video to upload": "image",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "时尚走秀.mp4",
"type": "input",
"format": "video/mp4",
"force_rate": 24,
"custom_width": 480,
"custom_height": 832,
"frame_load_cap": 121,
"skip_first_frames": 0,
"select_every_nth": 1
}
}
}
},
{
"id": 227,
"type": "easy sam3VideoSegmentation",
"pos": [
-4529.19864967506,
443.5309937256167
],
"size": [
314.00554584873043,
437.1614502609848
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "sam3_model",
"type": "EASY_SAM3_MODEL",
"link": 312
},
{
"name": "video_frames",
"type": "IMAGE",
"link": 313
},
{
"name": "session_id",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "extra_config",
"shape": 7,
"type": "EASY_SAM3_EXTRA_CONFIG",
"link": null
},
{
"name": "positive_coords",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "negative_coords",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "bbox",
"shape": 7,
"type": "BBOX",
"link": null
}
],
"outputs": [
{
"name": "masks",
"type": "MASK",
"links": [
318
]
},
{
"name": "session_id",
"type": "STRING",
"links": null
},
{
"name": "objects",
"type": "EASY_SAM3_OBJECTS_OUTPUT",
"links": null
}
],
"properties": {
"Node name for S&R": "easy sam3VideoSegmentation"
},
"widgets_values": [
"dress",
0,
1,
0.5,
0.7,
"both",
0,
-1,
true,
false
]
}
],
"links": [
[
312,
228,
0,
227,
0,
"EASY_SAM3_MODEL"
],
[
313,
226,
0,
227,
1,
"IMAGE"
],
[
318,
227,
0,
231,
0,
"MASK"
],
[
319,
231,
0,
232,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8579069603827805,
"offset": [
5332.017781339226,
351.4800928063147
]
},
"workflowRendererVersion": "LG",
"frontendVersion": "1.33.5",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
+66 -14
View File
@@ -51,6 +51,22 @@
"add_background": {
"name": "add_background",
"tooltip": "Add background color to segmented images"
},
"coordinates_positive": {
"name": "coordinates_positive",
"tooltip": "Positive click coordinates for refinement"
},
"coordinates_negative": {
"name": "coordinates_negative",
"tooltip": "Negative click coordinates for refinement"
},
"bboxes": {
"name": "bboxes",
"tooltip": "Bounding boxes for object detection"
},
"mask": {
"name": "mask",
"tooltip": "Input mask for refinement"
}
},
"outputs": {
@@ -61,6 +77,14 @@
"1": {
"name": "images",
"tooltip": "Segmentation images"
},
"2": {
"name": "boxes",
"tooltip": "Detected bounding boxes"
},
"3": {
"name": "scores",
"tooltip": "Detection confidence scores"
}
}
},
@@ -83,6 +107,14 @@
"name": "prompt",
"tooltip": "Text description of objects to track (e.g., 'person', 'car')"
},
"frame_index": {
"name": "frame_index",
"tooltip": "Frame where initial prompt is applied"
},
"object_id": {
"name": "object_id",
"tooltip": "Unique ID for multi-object tracking"
},
"score_threshold_detection": {
"name": "score_threshold_detection",
"tooltip": "Confidence threshold for detections, default is 0.5"
@@ -91,33 +123,41 @@
"name": "new_det_thresh",
"tooltip": "Threshold for a detection to be added as a new object, default is 0.7"
},
"object_frame_index": {
"name": "object_frame_index",
"tooltip": "Object Frame index to start tracking from"
},
"object_id": {
"name": "object_id",
"tooltip": "Object ID to track (1-100)"
},
"start_to_propagate": {
"name": "start_to_propagate",
"tooltip": "Propagation direction: disabled, both, forward, or backward"
"propagation_direction": {
"name": "propagation_direction",
"tooltip": "Propagation direction: both, forward, or backward"
},
"start_frame_index": {
"name": "start_frame_index",
"tooltip": "Frame index to start propagation from"
},
"keep_model_loaded": {
"name": "keep_model_loaded",
"tooltip": "Keep model in VRAM after inference"
"max_frames_to_track": {
"name": "max_frames_to_track",
"tooltip": "Advanced: Max frames to process (-1 for all)"
},
"close_after_propagation": {
"name": "close_after_propagation",
"tooltip": "Close the session after propagation"
},
"keep_model_loaded": {
"name": "keep_model_loaded",
"tooltip": "Keep model in VRAM after inference"
},
"extra_config": {
"name": "extra_config",
"tooltip": "Extra configuration for the SAM3 model"
},
"positive_coords": {
"name": "positive_coords",
"tooltip": "Positive click coordinates as JSON: '[{\"x\": 50, \"y\": 120}]'"
},
"negative_coords": {
"name": "negative_coords",
"tooltip": "Negative click coordinates as JSON: '[{\"x\": 150, \"y\": 300}]'"
},
"bbox": {
"name": "bbox",
"tooltip": "Bounding box as (x_min, y_min, x_max, y_max) or (x, y, width, height) tuple. Compatible with KJNodes Points Editor bbox output."
}
},
"outputs": {
@@ -128,6 +168,10 @@
"1": {
"name": "session_id",
"tooltip": "Session ID for resuming tracking"
},
"2": {
"name": "objects",
"tooltip": "Tracked objects output data"
}
}
},
@@ -174,6 +218,10 @@
"name": "det_nms_thresh",
"tooltip": "IoU threshold for detection NMS (Non-Maximum Suppression)"
},
"new_det_thresh": {
"name": "new_det_thresh",
"tooltip": "Threshold for a detection to be added as a new object"
},
"suppress_overlap_occlusion_thresh": {
"name": "suppress_overlap_occlusion_thresh",
"tooltip": "Threshold for suppressing overlapping objects based on recent occlusion (0.0 to disable)"
@@ -209,6 +257,10 @@
"masklet_confirm_frames": {
"name": "masklet_confirm_frames",
"tooltip": "Frames needed for masklet confirmation"
},
"image_size": {
"name": "image_size",
"tooltip": "Input image size for the model"
}
},
"outputs": {
+66 -14
View File
@@ -51,6 +51,22 @@
"add_background": {
"name": "添加背景",
"tooltip": "为分割的图像添加背景颜色"
},
"coordinates_positive": {
"name": "正向坐标",
"tooltip": "用于细化的正向点击坐标"
},
"coordinates_negative": {
"name": "负向坐标",
"tooltip": "用于细化的负向点击坐标"
},
"bboxes": {
"name": "边界框",
"tooltip": "用于对象检测的边界框"
},
"mask": {
"name": "遮罩",
"tooltip": "用于细化的输入遮罩"
}
},
"outputs": {
@@ -61,6 +77,14 @@
"1": {
"name": "图像",
"tooltip": "分割图像"
},
"2": {
"name": "边界框",
"tooltip": "检测到的边界框"
},
"3": {
"name": "分数",
"tooltip": "检测置信度分数"
}
}
},
@@ -83,6 +107,14 @@
"name": "提示词",
"tooltip": "要跟踪的对象的文本描述(例如:'人'、'汽车')"
},
"frame_index": {
"name": "帧索引",
"tooltip": "应用初始提示词的帧"
},
"object_id": {
"name": "对象 ID",
"tooltip": "多对象跟踪的唯一 ID"
},
"score_threshold_detection": {
"name": "检测分数阈值",
"tooltip": "检测的置信度阈值,默认为 0.5"
@@ -91,33 +123,41 @@
"name": "新检测阈值",
"tooltip": "将检测添加为新对象的阈值,默认为 0.7"
},
"object_frame_index": {
"name": "对象帧索引",
"tooltip": "开始跟踪的对象帧索引"
},
"object_id": {
"name": "对象 ID",
"tooltip": "要跟踪的对象 ID(1-100)"
},
"start_to_propagate": {
"name": "开始传播",
"tooltip": "传播方向:禁用、双向、前向或后向"
"propagation_direction": {
"name": "传播方向",
"tooltip": "传播方向:双向、前向或后向"
},
"start_frame_index": {
"name": "起始帧索引",
"tooltip": "开始传播的帧索引"
},
"keep_model_loaded": {
"name": "保持模型加载",
"tooltip": "推理后将模型保留在显存中"
"max_frames_to_track": {
"name": "最大跟踪帧数",
"tooltip": "高级:要处理的最大帧数(-1 表示全部)"
},
"close_after_propagation": {
"name": "传播后关闭",
"tooltip": "传播后关闭会话"
},
"keep_model_loaded": {
"name": "保持模型加载",
"tooltip": "推理后将模型保留在显存中"
},
"extra_config": {
"name": "额外配置",
"tooltip": "SAM3 模型的额外配置"
},
"positive_coords": {
"name": "正向坐标",
"tooltip": "正向点击坐标,JSON 格式:'[{\"x\": 50, \"y\": 120}]'"
},
"negative_coords": {
"name": "负向坐标",
"tooltip": "负向点击坐标,JSON 格式:'[{\"x\": 150, \"y\": 300}]'"
},
"bbox": {
"name": "边界框",
"tooltip": "边界框,格式为 (x_min, y_min, x_max, y_max) 或 (x, y, width, height) 元组。兼容 KJNodes Points Editor bbox 输出。"
}
},
"outputs": {
@@ -128,6 +168,10 @@
"1": {
"name": "会话 ID",
"tooltip": "用于恢复跟踪的会话 ID"
},
"2": {
"name": "对象",
"tooltip": "跟踪对象输出数据"
}
}
},
@@ -174,6 +218,10 @@
"name": "检测 NMS 阈值",
"tooltip": "检测 NMS(非极大值抑制)的 IoU 阈值"
},
"new_det_thresh": {
"name": "新检测阈值",
"tooltip": "将检测添加为新对象的阈值"
},
"suppress_overlap_occlusion_thresh": {
"name": "抑制重叠遮挡阈值",
"tooltip": "基于最近遮挡抑制重叠对象的阈值(0.0 禁用)"
@@ -209,6 +257,10 @@
"masklet_confirm_frames": {
"name": "掩码确认帧数",
"tooltip": "掩码确认所需的帧数"
},
"image_size": {
"name": "图像尺寸",
"tooltip": "模型的输入图像尺寸"
}
},
"outputs": {
+165 -122
View File
@@ -13,7 +13,7 @@ from PIL import Image
from typing import Tuple, Any
from comfy_api.latest import ComfyExtension, io
from .sam3.logger import get_logger
from .utils import tensor_to_pil, pil_to_tensor, masks_to_tensor, join_image_with_alpha
from .utils import tensor_to_pil, pil_to_tensor, masks_to_tensor, join_image_with_alpha, parse_points, parse_bbox
logger = get_logger(__name__)
@@ -40,6 +40,7 @@ class LoadSam3Model(io.ComfyNode):
io.Combo.Input(
"model",
options=folder_paths.get_filename_list("sam3"),
default="sam3.pt",
tooltip="Select SAM3 model file to load"
),
io.Combo.Input(
@@ -78,6 +79,9 @@ class LoadSam3Model(io.ComfyNode):
if model_path is None:
raise ValueError(f"Model file '{model}' not found in sam3 folder")
if "fp16" in model.lower():
precision = "fp16"
# Build model based on segmentor type
if segmentor == "image":
from .sam3.model.sam3_image_processor import Sam3Processor
@@ -107,6 +111,14 @@ class LoadSam3Model(io.ComfyNode):
logger.info("Sam3 Model loaded successfully")
if precision != 'fp32' and device == 'cpu':
raise ValueError("fp16 and bf16 are not supported on cpu")
if device == "cuda":
if torch.cuda.get_device_properties(0).major >= 8:
# turn on tfloat32 for Ampere GPUs (https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
device = {"cuda": torch.device("cuda"), "cpu": torch.device("cpu"), "mps": torch.device("mps")}[device]
@@ -149,7 +161,7 @@ class Sam3ImageSegmentation(io.ComfyNode):
),
io.Float.Input(
"threshold",
default=0.60,
default=0.40,
min=0.0,
max=1.0,
step=0.05,
@@ -164,33 +176,29 @@ class Sam3ImageSegmentation(io.ComfyNode):
options=["none", "black", "white", "grey"],
default="none",
tooltip="Add background color to segmented images"
)
# io.String.Input(
# "coordinates_positive",
# display_name="coordinates_positive",
# optional=True,
# force_input=True,
# ),
# io.String.Input(
# "coordinates_negative",
# display_name="coordinates_negative",
# optional=True,
# force_input=True,
# ),
# io.BBOX.Input(
# "bboxes",
# display_name="bboxes",
# optional=True,
# ),
# io.Mask.Input(
# "mask",
# display_name="mask",
# optional=True,
# ),
# io.Boolean.Input(
# "enable_visualize",
# default=False,
# ),
),
io.String.Input(
"coordinates_positive",
display_name="coordinates_positive",
optional=True,
force_input=True,
),
io.String.Input(
"coordinates_negative",
display_name="coordinates_negative",
optional=True,
force_input=True,
),
io.BBOX.Input(
"bboxes",
display_name="bboxes",
optional=True,
),
io.Mask.Input(
"mask",
display_name="mask",
optional=True,
),
],
outputs=[
io.Mask.Output(
@@ -205,12 +213,16 @@ class Sam3ImageSegmentation(io.ComfyNode):
is_output_list=True,
tooltip="Segmentation images",
),
# io.Image.Output(
# "visualization",
# display_name="visualization",
# is_output_list=True,
# tooltip="When enable_visualize is True, the visualized image is output, otherwise the original image is output.",
# )
io.String.Output(
"boxes",
display_name="boxes",
is_output_list=True,
),
io.String.Output(
"scores",
display_name="scores",
is_output_list=True,
),
]
)
@@ -232,56 +244,36 @@ class Sam3ImageSegmentation(io.ComfyNode):
# set confidence threshold
processor.set_confidence_threshold(threshold)
# Todo: support for points and bbox prompts
# # handle point coordinates
# if coordinates_positive is not None:
# try:
# coordinates_positive = json.loads(coordinates_positive.replace("'", '"'))
# coordinates_positive = [(coord['x'], coord['y']) for coord in coordinates_positive]
# if coordinates_negative is not None:
# coordinates_negative = json.loads(coordinates_negative.replace("'", '"'))
# coordinates_negative = [(coord['x'], coord['y']) for coord in coordinates_negative]
# except:
# pass
# Parse inputs with bounds checking
pos_points, pos_count, pos_errors = parse_points(coordinates_positive, images.shape)
neg_points, neg_count, neg_errors = parse_points(coordinates_negative, images.shape)
# Combine points for refinement
points = None
point_labels = None
if pos_points is not None and neg_points is not None:
points = pos_points + neg_points
point_labels = [1] * pos_count + [0] * neg_count
elif pos_points is not None:
points = pos_points
point_labels = [1] * pos_count
elif neg_points is not None:
points = neg_points
point_labels = [0] * neg_count
# positive_point_coords = np.atleast_2d(np.array(coordinates_positive))
# bbox
bounding_boxes = None
bounding_box_labels = None
if bboxes is not None:
bbox_coords, bbox_count = parse_bbox(bboxes, images.shape)
if bbox_coords is not None:
bounding_boxes = bbox_coords
bounding_box_labels = [1] * bbox_count
# if coordinates_negative is not None:
# negative_point_coords = np.array(coordinates_negative)
# # Ensure both positive and negative coords are lists of 2D arrays if individual_objects is True
# final_coords = np.concatenate((positive_point_coords, negative_point_coords), axis=0)
# else:
# final_coords = positive_point_coords
# # Handle possible bboxes
# if bboxes is not None:
# boxes_np_batch = []
# for bbox_list in bboxes:
# boxes_np = []
# for bbox in bbox_list:
# boxes_np.append(bbox)
# boxes_np = np.array(boxes_np)
# boxes_np_batch.append(boxes_np)
# final_box = np.array(boxes_np)
# final_labels = None
# # handle labels
# if coordinates_positive is not None:
# positive_point_labels = np.ones(len(positive_point_coords))
# if coordinates_negative is not None:
# negative_point_labels = np.zeros(len(negative_point_coords)) # 0 = negative
# final_labels = np.concatenate((positive_point_labels, negative_point_labels), axis=0)
# else:
# final_labels = positive_point_labels
# print("combined labels: ", final_labels)
# print("combined labels shape: ", final_labels.shape)
# mask_list = []
# Switch model to main device
model.to(device)
if mask is not None:
mask.to(device)
autocast_condition = not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
@@ -293,6 +285,8 @@ class Sam3ImageSegmentation(io.ComfyNode):
# Process each image separately to maintain correspondence
output_masks = []
output_images = []
output_boxes = []
output_scores = []
# Initialize progress bar
@@ -308,7 +302,15 @@ class Sam3ImageSegmentation(io.ComfyNode):
if prompt_text:
state = processor.set_text_prompt(prompt_text, state)
# TODO: Add support for points and bbox prompts
# points
if points is not None and len(points) > 0:
state = processor.add_point_prompt(points, point_labels, state)
# bbox
if bounding_boxes is not None and len(bounding_boxes) > 0:
state = processor.add_multiple_box_prompts(bounding_boxes, bounding_box_labels, state)
# mask
if mask is not None:
state = processor.add_mask_prompt(mask, state)
# Get the masks and scores for this image
masks = state.get('masks', None)
@@ -345,7 +347,7 @@ class Sam3ImageSegmentation(io.ComfyNode):
img_tensor = pil_to_tensor(pil_img)
mask_tensor = combined_mask.unsqueeze(0)
rgba_image, = join_image_with_alpha(img_tensor, mask_tensor, False)
if add_background != "none":
if add_background == "black":
bg_color = torch.zeros_like(rgba_image[:, :, :, :3])
@@ -353,26 +355,17 @@ class Sam3ImageSegmentation(io.ComfyNode):
bg_color = torch.ones_like(rgba_image[:, :, :, :3])
elif add_background == "grey":
bg_color = torch.ones_like(rgba_image[:, :, :, :3]) * 0.5
rgb = rgba_image[:, :, :, :3]
alpha = rgba_image[:, :, :, 3:4]
composited = rgb * alpha + bg_color * (1 - alpha)
output_images.append([composited.squeeze(0)])
else:
output_images.append([rgba_image.squeeze(0)])
# Visualization: overlay mask on original image with color
# if enable_visualize:
# vis_image = visualize_masks_on_image(
# pil_img,
# combined_mask,
# boxes,
# scores,
# alpha=0.5
# )
# vis_tensor = pil_to_tensor(vis_image)
# output_visualizations.append(vis_tensor)
output_boxes.append(boxes)
output_scores.append(scores)
# Update progress bar
processed_frames += 1
@@ -386,7 +379,7 @@ class Sam3ImageSegmentation(io.ComfyNode):
model.to(offload_device)
mm.soft_empty_cache()
return io.NodeOutput(output_masks, output_images)
return io.NodeOutput(output_masks, output_images, output_boxes, output_scores)
class Sam3VideoSegmentation(io.ComfyNode):
@@ -425,6 +418,15 @@ class Sam3VideoSegmentation(io.ComfyNode):
min=0,
max=10 ** 5,
step=1,
tooltip="Frame where initial prompt is applied",
),
io.Int.Input(
"object_id",
default=1,
min=1,
max=1000,
step=1,
tooltip="Unique ID for multi-object tracking"
),
io.Float.Input(
"score_threshold_detection",
@@ -454,6 +456,12 @@ class Sam3VideoSegmentation(io.ComfyNode):
max=10**5,
step=1,
),
io.Int.Input(
"max_frames_to_track",
default=-1,
min=-1,
tooltip="Advanced: Max frames to process (-1 for all)"
),
io.Boolean.Input(
"close_after_propagation",
default=True,
@@ -469,28 +477,26 @@ class Sam3VideoSegmentation(io.ComfyNode):
tooltip="Extra configuration for the SAM3 model",
optional=True,
),
# io.String.Input(
# "coordinates_positive",
# display_name="coordinates_positive",
# optional=True,
# force_input=True,
# ),
# io.String.Input(
# "coordinates_negative",
# display_name="coordinates_negative",
# optional=True,
# force_input=True,
# ),
# io.BBOX.Input(
# "bboxes",
# display_name="bboxes",
# optional=True,
# ),
# io.Mask.Input(
# "mask",
# display_name="mask",
# optional=True,
# )
io.String.Input(
"positive_coords",
display_name="positive_coords",
tooltip="Positive click coordinates as JSON: '[{\"x\": 50, \"y\": 120}]'",
optional=True,
force_input=True,
),
io.String.Input(
"negative_coords",
display_name="negative_coords",
tooltip="Negative click coordinates as JSON: '[{\"x\": 150, \"y\": 300}]'",
optional=True,
force_input=True,
),
io.BBOX.Input(
"bbox",
display_name="bbox",
optional=True,
tooltip="Bounding box as (x_min, y_min, x_max, y_max) or (x, y, width, height) tuple. Compatible with KJNodes Points Editor bbox output."
),
],
outputs=[
io.Mask.Output(
@@ -511,8 +517,8 @@ class Sam3VideoSegmentation(io.ComfyNode):
@classmethod
def execute(cls, sam3_model, video_frames, prompt, frame_index, score_threshold_detection, new_det_thresh, propagation_direction, start_frame_index=0,close_after_propagation=True, keep_model_loaded=False, session_id=None, extra_config=None,coordinates_positive=None, coordinates_negative=None,
bboxes=None, mask=None) -> io.NodeOutput:
def execute(cls, sam3_model, video_frames, prompt, frame_index, object_id, score_threshold_detection, new_det_thresh, propagation_direction, start_frame_index=0, max_frames_to_track=-1, close_after_propagation=True, keep_model_loaded=False, session_id=None, extra_config=None, positive_coords=None, negative_coords=None,
bbox=None,) -> io.NodeOutput:
offload_device = mm.unet_offload_device()
video_predictor = sam3_model.get("model", None)
@@ -524,6 +530,10 @@ class Sam3VideoSegmentation(io.ComfyNode):
if video_predictor is None or segmentor != "video":
raise ValueError("Invalid SAM3 model. Please load a SAM3 model in 'video' mode")
if frame_index > B - 1:
logger.info(f"Frame index {frame_index} is out of bounds, setting to last frame {B - 1}")
frame_index = B - 1
# Set video model config
video_predictor.model.score_threshold_detection = score_threshold_detection
video_predictor.model.new_det_thresh = new_det_thresh
@@ -574,14 +584,46 @@ class Sam3VideoSegmentation(io.ComfyNode):
video_predictor.model.to(device)
autocast_condition = not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
# Parse inputs with bounds checking
pos_points, pos_count, pos_errors = parse_points(positive_coords, video_frames.shape)
neg_points, neg_count, neg_errors = parse_points(negative_coords, video_frames.shape)
# Combine points for refinement
points = None
point_labels = None
if pos_points is not None and neg_points is not None:
points = pos_points + neg_points
point_labels = [1] * pos_count + [0] * neg_count
elif pos_points is not None:
points = pos_points
point_labels = [1] * pos_count
elif neg_points is not None:
points = neg_points
point_labels = [0] * neg_count
# bbox (has bugs)
bounding_boxes = None
bounding_box_labels = None
if bbox is not None:
bbox_coords, bbox_count = parse_bbox(bbox, video_frames.shape)
if bbox_coords is not None:
bounding_boxes = bbox_coords
bounding_box_labels = [1] * bbox_count
# print('bbox_coords:', bbox_coords)
# Add Prompt
response = video_predictor.handle_request(
request=dict(
type="add_prompt",
session_id=session_id,
frame_index=frame_index,
text=prompt,
text=prompt if prompt else None,
bounding_boxes=bounding_boxes,
bounding_box_labels=bounding_box_labels,
points=points,
point_labels=point_labels,
obj_id=object_id
)
)
@@ -603,6 +645,7 @@ class Sam3VideoSegmentation(io.ComfyNode):
session_id=session_id,
propagation_direction=propagation_direction,
start_frame_index=start_frame_index,
max_frame_num_to_track=max_frames_to_track if max_frames_to_track != -1 else None,
)
):
frame_idx = response.get("frame_index", 0)
+84
View File
@@ -151,6 +151,90 @@ class Sam3Processor:
return self._forward_grounding(state)
@torch.inference_mode()
def add_multiple_box_prompts(self, boxes: List[List], labels: List[bool], state: Dict):
"""Adds multiple box prompts and run the inference.
The image needs to be set, but not necessarily the text prompt.
Each box is assumed to be in [center_x, center_y, width, height] format and normalized in [0, 1] range.
Each label is True for a positive box, False for a negative box.
"""
if "backbone_out" not in state:
raise ValueError("You must call set_image before add_multiple_box_prompts")
if "language_features" not in state["backbone_out"]:
dummy_text_outputs = self.model.backbone.forward_text(
["visual"], device=self.device
)
state["backbone_out"].update(dummy_text_outputs)
if "geometric_prompt" not in state:
state["geometric_prompt"] = self.model._get_dummy_prompt()
# Convert to [seq_len, batch_size, 4] format
boxes_tensor = torch.tensor(boxes, device=self.device, dtype=torch.float32).view(len(boxes), 1, 4)
labels_tensor = torch.tensor(labels, device=self.device, dtype=torch.bool).view(len(labels), 1)
state["geometric_prompt"].append_boxes(boxes_tensor, labels_tensor)
return self._forward_grounding(state)
@torch.inference_mode()
def add_point_prompt(self, points: List[List], labels: List[int], state: Dict):
"""Adds point prompts and run the inference.
The image needs to be set, but not necessarily the text prompt.
Points should be in [x, y] format, normalized in [0, 1] range.
Labels should be 1 for foreground points, 0 for background points.
"""
if "backbone_out" not in state:
raise ValueError("You must call set_image before add_point_prompt")
if "language_features" not in state["backbone_out"]:
dummy_text_outputs = self.model.backbone.forward_text(
["visual"], device=self.device
)
state["backbone_out"].update(dummy_text_outputs)
if "geometric_prompt" not in state:
state["geometric_prompt"] = self.model._get_dummy_prompt()
# Convert to [seq_len, batch_size, 2] format
points_tensor = torch.tensor(points, device=self.device, dtype=torch.float32).view(len(points), 1, 2)
labels_tensor = torch.tensor(labels, device=self.device, dtype=torch.long).view(len(labels), 1)
state["geometric_prompt"].append_points(points_tensor, labels_tensor)
return self._forward_grounding(state)
@torch.inference_mode()
def add_mask_prompt(self, mask: torch.Tensor, state: Dict):
"""Adds a mask prompt and run the inference.
The mask should be a binary tensor with shape matching the model's expected input.
This is typically used for iterative refinement.
"""
if "backbone_out" not in state:
raise ValueError("You must call set_image before add_mask_prompt")
if "language_features" not in state["backbone_out"]:
dummy_text_outputs = self.model.backbone.forward_text(
["visual"], device=self.device
)
state["backbone_out"].update(dummy_text_outputs)
if "geometric_prompt" not in state:
state["geometric_prompt"] = self.model._get_dummy_prompt()
# Ensure mask is on correct device and has batch dimension
if mask.device != self.device:
mask = mask.to(self.device)
# Add sequence and batch dimensions if needed: [seq_len, batch_size, H, W]
if len(mask.shape) == 2: # [H, W]
mask = mask.unsqueeze(0).unsqueeze(0)
elif len(mask.shape) == 3: # [1, H, W] or [batch, H, W]
mask = mask.unsqueeze(0)
state["geometric_prompt"].append_masks(mask)
return self._forward_grounding(state)
def reset_all_prompts(self, state: Dict):
"""Removes all the prompts and results"""
if "backbone_out" in state:
+1 -1
View File
@@ -482,7 +482,7 @@ class Sam3TrackerPredictor(Sam3TrackerBase):
if self.non_overlap_masks_for_output:
video_res_masks = self._apply_non_overlapping_constraints(video_res_masks)
# potentially fill holes in the predicted masks
if self.fill_hole_area > 0:
if self.fill_hole_area > 0 and len(video_res_masks) > 0:
video_res_masks = fill_holes_in_mask_scores(
video_res_masks, self.fill_hole_area
)
+7 -6
View File
@@ -966,12 +966,13 @@ class Sam3VideoBase(nn.Module):
# Part 2: masks from new detections
new_det_fa_inds_t = torch.from_numpy(new_det_fa_inds)
new_det_low_res_masks = det_out["mask"][new_det_fa_inds_t].unsqueeze(1)
new_det_low_res_masks = fill_holes_in_mask_scores(
new_det_low_res_masks,
max_area=self.fill_hole_area,
fill_holes=True,
remove_sprinkles=True,
)
if len(new_det_fa_inds) > 0:
new_det_low_res_masks = fill_holes_in_mask_scores(
new_det_low_res_masks,
max_area=self.fill_hole_area,
fill_holes=True,
remove_sprinkles=True,
)
new_masklet_video_res_masks = F.interpolate(
new_det_low_res_masks,
size=(orig_vid_height, orig_vid_width),
+14 -2
View File
@@ -84,7 +84,7 @@ class Sam3VideoInference(Sam3VideoBase):
inference_state["feature_cache"] = {}
inference_state["cached_frame_outputs"] = {}
inference_state["action_history"] = [] # for logging user actions
inference_state["is_image_only"] = is_image_type(resource_path) if resource_path else False
inference_state["is_image_only"] = is_image_type(resource_path)
return inference_state
@torch.inference_mode()
@@ -1154,6 +1154,12 @@ class Sam3VideoInferenceWithInstanceInteractivity(Sam3VideoInference):
) # (1, H_video, W_video) bool
refined_obj_id_to_mask[obj_id] = refined_mask_video_res
# Initialize cache if not present (needed for point prompts during propagation)
if "cached_frame_outputs" not in inference_state:
inference_state["cached_frame_outputs"] = {}
if frame_idx not in inference_state["cached_frame_outputs"]:
inference_state["cached_frame_outputs"][frame_idx] = {}
obj_id_to_mask = self._build_tracker_output(
inference_state, frame_idx, refined_obj_id_to_mask
)
@@ -1575,6 +1581,12 @@ class Sam3VideoInferenceWithInstanceInteractivity(Sam3VideoInference):
new_mask_data = data_list[0].to(self.device)
if self.rank == 0:
# Initialize cache if not present (needed for point prompts without prior propagation)
if "cached_frame_outputs" not in inference_state:
inference_state["cached_frame_outputs"] = {}
if frame_idx not in inference_state["cached_frame_outputs"]:
inference_state["cached_frame_outputs"][frame_idx] = {}
obj_id_to_mask = self._build_tracker_output(
inference_state,
frame_idx,
@@ -1706,4 +1718,4 @@ class Sam3VideoInferenceWithInstanceInteractivity(Sam3VideoInference):
def is_image_type(resource_path: str) -> bool:
if isinstance(resource_path, list):
return len(resource_path) == 1
return resource_path.lower().endswith(tuple(IMAGE_EXTS))
return resource_path.lower().endswith(tuple(IMAGE_EXTS))
-12
View File
@@ -309,18 +309,6 @@ def _create_sam3_model(
}
matcher = None
if not eval_mode:
from .train.matcher import BinaryHungarianMatcherV2
matcher = BinaryHungarianMatcherV2(
focal=True,
cost_class=2.0,
cost_bbox=5.0,
cost_giou=2.0,
alpha=0.25,
gamma=2,
stable=False,
)
common_params["matcher"] = matcher
model = Sam3Image(**common_params)
+281
View File
@@ -0,0 +1,281 @@
{
"id": "307743a2-8a4b-4090-b8dc-3125ff31fa14",
"revision": 0,
"last_node_id": 223,
"last_link_id": 304,
"nodes": [
{
"id": 222,
"type": "PreviewImage",
"pos": [
-3883.9588057841556,
593.0894010290336
],
"size": [
339.2921196124935,
246
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 303
}
],
"outputs": [],
"properties": {
"cnr_id": "comfy-core",
"ver": "0.3.70",
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 218,
"type": "MaskPreview+",
"pos": [
-3880.8437110898844,
280.9015183807766
],
"size": [
332.95740427927046,
258
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "mask",
"type": "MASK",
"link": 298
}
],
"outputs": [],
"properties": {
"cnr_id": "comfyui_essentials",
"ver": "9d9f4bedfc9f0321c19faf71855e228c93bd0dc9",
"Node name for S&R": "MaskPreview+"
},
"widgets_values": []
},
{
"id": 190,
"type": "easy sam3ModelLoader",
"pos": [
-4235.041825672723,
334.7519653807837
],
"size": [
314.09852906952574,
130
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "sam3_model",
"type": "EASY_SAM3_MODEL",
"links": [
282
]
}
],
"properties": {
"Node name for S&R": "easy sam3ModelLoader"
},
"widgets_values": [
"sam3.pt",
"image",
"cuda",
"fp16"
]
},
{
"id": 221,
"type": "LoadImage",
"pos": [
-4678.319118164902,
372.4935004943305
],
"size": [
406.60579628243977,
446.341753250938
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
302
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"cnr_id": "comfy-core",
"ver": "0.3.70",
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"1 (3).png",
"image"
]
},
{
"id": 213,
"type": "easy sam3ImageSegmentation",
"pos": [
-4241.387303264146,
538.9760914950197
],
"size": [
337.2438563422779,
281.966544722883
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "sam3_model",
"type": "EASY_SAM3_MODEL",
"link": 282
},
{
"name": "images",
"type": "IMAGE",
"link": 302
},
{
"name": "coordinates_positive",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "coordinates_negative",
"shape": 7,
"type": "STRING",
"link": null
},
{
"name": "bboxes",
"shape": 7,
"type": "BBOX",
"link": null
},
{
"name": "mask",
"shape": 7,
"type": "MASK",
"link": null
}
],
"outputs": [
{
"name": "masks",
"shape": 6,
"type": "MASK",
"links": [
298
]
},
{
"name": "images",
"shape": 6,
"type": "IMAGE",
"links": [
303
]
},
{
"name": "boxes",
"shape": 6,
"type": "STRING",
"links": null
},
{
"name": "scores",
"shape": 6,
"type": "STRING",
"links": null
}
],
"properties": {
"Node name for S&R": "easy sam3ImageSegmentation"
},
"widgets_values": [
"apple",
0.4,
false,
"none"
]
}
],
"links": [
[
282,
190,
0,
213,
0,
"EASY_SAM3_MODEL"
],
[
298,
213,
0,
218,
0,
"MASK"
],
[
302,
221,
0,
213,
1,
"IMAGE"
],
[
303,
213,
1,
222,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 1.3920191078517683,
"offset": [
4898.859436279234,
-79.3314399803767
]
},
"workflowRendererVersion": "LG",
"frontendVersion": "1.33.5",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
+214 -1
View File
@@ -4,6 +4,7 @@ Utility functions for tensor and PIL image conversions.
import numpy as np
import torch
import json
from PIL import Image
from typing import List, Union, Optional
@@ -228,4 +229,216 @@ def join_image_with_alpha(image: torch.Tensor, alpha: torch.Tensor, invert=False
for i in range(batch_size):
out_images.append(torch.cat((image[i][:,:,:3], alpha[i].unsqueeze(2)), dim=2))
return torch.stack(out_images),
return torch.stack(out_images),
def parse_points(points_str, image_shape=None):
"""Parse point coordinates from JSON string and validate bounds.
Converts pixel coordinates to normalized coordinates (0-1 range) if image_shape is provided.
Returns:
tuple: (points_array, labels_array, validation_errors) where validation_errors
is a list of error messages, or (None, None, errors) if all points invalid
"""
if not points_str or not points_str.strip():
return None, None, []
try:
points_list = json.loads(points_str)
if not isinstance(points_list, list):
raise ValueError(f"Points must be a JSON array, got {type(points_list).__name__}")
if len(points_list) == 0:
return None, None, []
points = []
validation_errors = []
for i, point_dict in enumerate(points_list):
if not isinstance(point_dict, dict):
err = f"Point {i} is not a dictionary"
print(f"Warning: {err}, skipping")
validation_errors.append(err)
continue
if 'x' not in point_dict or 'y' not in point_dict:
err = f"Point {i} missing 'x' or 'y' key"
print(f"Warning: {err}, skipping")
validation_errors.append(err)
continue
try:
x = float(point_dict['x'])
y = float(point_dict['y'])
# Validate coordinates are non-negative
if x < 0 or y < 0:
err = f"Point {i} has negative coordinates ({x}, {y})"
print(f"Warning: {err}, skipping")
validation_errors.append(err)
continue
# Normalize to 0-1 range if image shape is provided
if image_shape is not None:
height, width = image_shape[1], image_shape[2] # [batch, height, width, channels]
# Validate within image bounds
if x >= width or y >= height:
err = f"Point {i} ({x}, {y}) outside image bounds ({width}x{height})"
print(f"Warning: {err}, skipping")
validation_errors.append(err)
continue
# Normalize coordinates to [0, 1] range
x = x / width
y = y / height
points.append([x, y])
except (ValueError, TypeError) as e:
err = f"Could not convert point {i} coordinates to float: {e}"
print(f"Warning: {err}, skipping")
validation_errors.append(err)
continue
if not points:
return None, None, validation_errors
return points, len(points), validation_errors
except json.JSONDecodeError as e:
raise ValueError(f"Invalid JSON in points: {str(e)}")
except Exception as e:
print(f"Error parsing points: {e}")
return None, None, [str(e)]
def parse_bbox(bbox, image_shape=None):
"""Parse bounding box from BBOX type (tuple/list/dict) and validate
Converts pixel coordinates to normalized coordinates (0-1 range) if image_shape is provided.
Supports multiple formats:
- KJNodes: [{'startX': x, 'startY': y, 'endX': x2, 'endY': y2}, ...]
- Tuple/list: (x1, y1, x2, y2) or (x, y, width, height)
- Dict: {'startX': x, 'startY': y, 'endX': x2, 'endY': y2}
Returns:
List of bounding boxes [[x1, y1, x2, y2], ...] in normalized coordinates (0-1) if image_shape provided, or None
"""
if bbox is None:
return None
try:
all_coords = []
# Try to extract coordinates regardless of type checks
# This handles cases where ComfyUI wraps data in unexpected ways
if hasattr(bbox, '__iter__') and not isinstance(bbox, (str, bytes)):
# It's some kind of sequence
try:
bbox_list = list(bbox)
if len(bbox_list) == 0:
return None
# Check if it's a list of 4 numbers (single bbox)
if len(bbox_list) == 4 and all(isinstance(x, (int, float)) for x in bbox_list):
coords = [float(x) for x in bbox_list]
all_coords.append(coords)
else:
# Process each element as a potential bbox
for elem in bbox_list:
coords = None
# Try to access as dict-like (KJNodes format)
if hasattr(elem, '__getitem__'):
try:
x1 = float(elem['startX'])
y1 = float(elem['startY'])
x2 = float(elem['endX'])
y2 = float(elem['endY'])
coords = [x1, y1, x2, y2]
except (KeyError, TypeError):
# Not dict format, might be numeric sequence
pass
# If still no coords, try as numeric sequence
if coords is None:
if hasattr(elem, '__iter__') and not isinstance(elem, (str, bytes)):
inner = list(elem)
if len(inner) == 4:
coords = [float(x) for x in inner]
if coords is not None:
all_coords.append(coords)
except Exception as e:
raise ValueError(f"Failed to process bbox as sequence: {e}")
# Try single dict format
elif hasattr(bbox, '__getitem__'):
try:
x1 = float(bbox['startX'])
y1 = float(bbox['startY'])
x2 = float(bbox['endX'])
y2 = float(bbox['endY'])
coords = [x1, y1, x2, y2]
all_coords.append(coords)
except (KeyError, TypeError) as e:
raise ValueError(f"Dictionary bbox missing required keys: {e}")
else:
raise ValueError(f"Unsupported bbox type: {type(bbox)}")
if not all_coords:
raise ValueError(
f"Could not extract coordinates from bbox. Type: {type(bbox)}, Content: {repr(bbox)[:200]}")
# Process and validate each bbox
validated_coords = []
for coords in all_coords:
# Handle xywh format (convert to xyxy)
x1, y1, x2, y2 = coords
if x2 < x1 or y2 < y1:
# Assume xywh format: (x, y, width, height)
width, height = x2, y2
x2 = x1 + width
y2 = y1 + height
coords = [x1, y1, x2, y2]
# Validate coordinates
if coords[0] >= coords[2]:
raise ValueError(f"Invalid bbox: x1 ({coords[0]}) must be < x2 ({coords[2]})")
if coords[1] >= coords[3]:
raise ValueError(f"Invalid bbox: y1 ({coords[1]}) must be < y2 ({coords[3]})")
if coords[0] < 0 or coords[1] < 0:
raise ValueError(f"Bounding box coordinates must be non-negative, got x1={coords[0]}, y1={coords[1]}")
# Normalize to 0-1 range if image shape is provided
if image_shape is not None:
height, width = image_shape[1], image_shape[2] # [batch, height, width, channels]
# Validate within image bounds
if coords[0] >= width or coords[2] > width:
print(f"Warning: bbox x coordinates ({coords[0]}, {coords[2]}) outside image width ({width})")
if coords[1] >= height or coords[3] > height:
print(f"Warning: bbox y coordinates ({coords[1]}, {coords[3]}) outside image height ({height})")
# Normalize coordinates to [0, 1] range
coords = [
coords[0] / width, # x1
coords[1] / height, # y1
coords[2] / width, # x2
coords[3] / height # y2
]
validated_coords.append(coords)
return validated_coords, len(validated_coords)
except (ValueError, TypeError) as e:
error_msg = f"Invalid bbox: {str(e)}\n"
error_msg += f"Input type: {type(bbox)}\n"
error_msg += f"Input content: {repr(bbox)[:500]}"
raise ValueError(error_msg)