Add detection_limit to image seg

This commit is contained in:
yolain
2025-11-24 02:19:37 +08:00
parent 3ac5e4d32c
commit 2d99efacdb
2 changed files with 33 additions and 17 deletions
+20 -16
View File
@@ -29,65 +29,69 @@
"display_name": "SAM3 图像分割",
"inputs": {
"sam3_model": {
"name": "sam3_model",
"name": "SAM3 模型",
"tooltip": "从 LoadSam3Model 节点加载的 SAM3 模型(必须是 'image' 模式)"
},
"images": {
"name": "images",
"name": "图像",
"tooltip": "输入要分割的图像"
},
"prompt": {
"name": "prompt",
"name": "提示词",
"tooltip": "要分割的对象的文本描述(例如:'一只猫'、'人')。支持空字符串以仅使用点/框进行分割"
},
"threshold": {
"name": "threshold",
"name": "置信度",
"tooltip": "检测的置信度阈值(0.0-1.0,默认:0.40)"
},
"keep_model_loaded": {
"name": "keep_model_loaded",
"name": "保持模型加载",
"tooltip": "推理后将模型保留在显存中(默认:False)"
},
"add_background": {
"name": "add_background",
"name": "添加背景",
"tooltip": "为分割的图像添加背景颜色(none、black、white、grey)"
},
"coordinates_positive": {
"name": "coordinates_positive",
"name": "正向点坐标",
"tooltip": "正向点坐标,JSON 字符串格式:'[{\"x\": 50, \"y\": 120}]'"
},
"coordinates_negative": {
"name": "coordinates_negative",
"name": "负向点坐标",
"tooltip": "负向点坐标,JSON 字符串格式:'[{\"x\": 150, \"y\": 300}]'"
},
"bboxes": {
"name": "bboxes",
"name": "边界框",
"tooltip": "边界框,格式为 (x_min, y_min, x_max, y_max) 或 (x, y, width, height)"
},
"mask": {
"name": "mask",
"name": "遮罩",
"tooltip": "用于细化的输入遮罩"
},
"detection_limit": {
"name": "检测对象限制",
"tooltip": "限制每张图像检测的最大对象数量(默认:-1)"
}
},
"outputs": {
"0": {
"name": "masks",
"name": "遮罩",
"tooltip": "合并的分割遮罩(每张图像一个遮罩,所有检测到的对象合并在一起)"
},
"1": {
"name": "images",
"name": "图像",
"tooltip": "带 RGBA 透明通道的分割图像(可选背景)"
},
"2": {
"name": "obj_masks",
"name": "对象这种",
"tooltip": "合并前的单个对象遮罩(用于可视化)"
},
"3": {
"name": "boxes",
"name": "边界框",
"tooltip": "每个检测对象的边界框坐标 [N, 4] 格式"
},
"4": {
"name": "scores",
"name": "置信度分",
"tooltip": "每个检测对象的置信度分数"
}
}
@@ -299,7 +303,7 @@
},
"outputs": {
"0": {
"name": "bbox",
"name": "边界框",
"tooltip": "解析后的 BBOX 格式边界框"
}
}
+13 -1
View File
@@ -199,6 +199,13 @@ class Sam3ImageSegmentation(io.ComfyNode):
display_name="mask",
optional=True,
),
io.Int.Input(
"detection_limit",
default=-1,
min=-1,
max=1000,
tooltip="Advanced: Limit number of detections (-1 for no limit)"
)
],
outputs=[
io.Mask.Output(
@@ -231,7 +238,7 @@ class Sam3ImageSegmentation(io.ComfyNode):
)
@classmethod
def execute(cls, sam3_model, images, prompt, threshold=0.3, keep_model_loaded=False, add_background='none', enable_visualize=False, coordinates_positive=None, coordinates_negative=None, bboxes=None, mask=None) -> io.NodeOutput:
def execute(cls, sam3_model, images, prompt, threshold=0.3, keep_model_loaded=False, add_background='none', detection_limit=-1, coordinates_positive=None, coordinates_negative=None, bboxes=None, mask=None) -> io.NodeOutput:
offload_device = mm.unet_offload_device()
processor = sam3_model.get("processor", None)
@@ -336,6 +343,11 @@ class Sam3ImageSegmentation(io.ComfyNode):
boxes = boxes[top_indices]
scores = scores[top_indices]
if detection_limit > -1:
masks = masks[:detection_limit]
boxes = boxes[:detection_limit]
scores = scores[:detection_limit]
output_raw_masks.append(masks)
# Convert masks to tensor format
masks_tensor = masks_to_tensor(masks)