diff --git a/comfy_extras/nodes_mask.py b/comfy_extras/nodes_mask.py index 3fae7221f..530782cdb 100644 --- a/comfy_extras/nodes_mask.py +++ b/comfy_extras/nodes_mask.py @@ -262,7 +262,7 @@ class MaskComposite(IO.ComfyNode): IO.Mask.Input("source"), IO.Int.Input("x", default=0, min=0, max=nodes.MAX_RESOLUTION, step=1), IO.Int.Input("y", default=0, min=0, max=nodes.MAX_RESOLUTION, step=1), - IO.Combo.Input("operation", options=["multiply", "add", "subtract", "and", "or", "xor"]), + IO.Combo.Input("operation", options=["multiply", "add", "subtract", "and", "or", "xor", "max", "min"], tooltip="How to combine the two masks. \"multiply\", \"add\", \"subtract\", \"max\" and \"min\" are arithmetic and preserve intermediate (feathered) mask values. \"and\", \"or\" and \"xor\" are boolean: they round each mask to 0 or 1 first, discarding any feathering. Use \"max\"/\"min\" for union/intersection of soft masks."), ], outputs=[IO.Mask.Output()], ) @@ -292,6 +292,10 @@ class MaskComposite(IO.ComfyNode): output[:, top:bottom, left:right] = torch.bitwise_or(destination_portion.round().bool(), source_portion.round().bool()).float() elif operation == "xor": output[:, top:bottom, left:right] = torch.bitwise_xor(destination_portion.round().bool(), source_portion.round().bool()).float() + elif operation == "max": + output[:, top:bottom, left:right] = torch.max(destination_portion, source_portion) + elif operation == "min": + output[:, top:bottom, left:right] = torch.min(destination_portion, source_portion) output = torch.clamp(output, 0.0, 1.0) diff --git a/tests-unit/comfy_extras_test/nodes_mask_test.py b/tests-unit/comfy_extras_test/nodes_mask_test.py new file mode 100644 index 000000000..a02032c3a --- /dev/null +++ b/tests-unit/comfy_extras_test/nodes_mask_test.py @@ -0,0 +1,58 @@ +import pytest +import torch +from unittest.mock import patch, MagicMock + +mock_nodes = MagicMock() +mock_nodes.MAX_RESOLUTION = 16384 +mock_server = MagicMock() + +with patch.dict("sys.modules", {"nodes": mock_nodes, "server": mock_server}): + from comfy_extras.nodes_mask import MaskComposite + + +class TestMaskCompositeOperations: + @staticmethod + def _exec(destination, source, operation): + d = torch.tensor([[destination]]) + s = torch.tensor([[source]]) + return MaskComposite.execute(d, s, 0, 0, operation).result[0].flatten().tolist() + + def test_max_is_union_of_soft_masks(self): + result = self._exec([0.0, 0.25, 0.75, 1.0], [0.5, 0.5, 0.5, 0.5], "max") + assert result == pytest.approx([0.5, 0.5, 0.75, 1.0]) + + def test_min_is_intersection_of_soft_masks(self): + result = self._exec([0.0, 0.25, 0.75, 1.0], [0.5, 0.5, 0.5, 0.5], "min") + assert result == pytest.approx([0.0, 0.25, 0.5, 0.5]) + + def test_max_preserves_intermediate_values(self): + # Unlike "or", max must not round feathered values to 0 or 1. + result = self._exec([0.25, 0.75], [0.0, 0.0], "max") + assert result == pytest.approx([0.25, 0.75]) + + def test_min_preserves_intermediate_values(self): + result = self._exec([0.25, 0.75], [1.0, 1.0], "min") + assert result == pytest.approx([0.25, 0.75]) + + def test_or_binarizes(self): + # Documents existing behaviour that motivates max/min. + result = self._exec([0.25, 0.75], [0.0, 0.0], "or") + assert result == pytest.approx([0.0, 1.0]) + + def test_max_is_commutative(self): + a = self._exec([0.3, 0.8], [0.6, 0.1], "max") + b = self._exec([0.6, 0.1], [0.3, 0.8], "max") + assert a == pytest.approx(b) + + def test_min_is_commutative(self): + a = self._exec([0.3, 0.8], [0.6, 0.1], "min") + b = self._exec([0.6, 0.1], [0.3, 0.8], "min") + assert a == pytest.approx(b) + + def test_max_with_empty_mask_is_identity(self): + result = self._exec([0.0, 0.4, 1.0], [0.0, 0.0, 0.0], "max") + assert result == pytest.approx([0.0, 0.4, 1.0]) + + def test_min_with_full_mask_is_identity(self): + result = self._exec([0.0, 0.4, 1.0], [1.0, 1.0, 1.0], "min") + assert result == pytest.approx([0.0, 0.4, 1.0])