diff --git a/comfy_api/latest/_io.py b/comfy_api/latest/_io.py index cdbffcb82..5ce99a49e 100644 --- a/comfy_api/latest/_io.py +++ b/comfy_api/latest/_io.py @@ -853,6 +853,7 @@ class Compositor(ComfyTypeIO): canvas: dict background: NotRequired[dict] inputs: NotRequired[list[str]] + order: NotRequired[list[int]] layers: list[dict] Type = LayerState diff --git a/comfy_extras/nodes_compositor.py b/comfy_extras/nodes_compositor.py index 2ea2cc05b..6afc12e4b 100644 --- a/comfy_extras/nodes_compositor.py +++ b/comfy_extras/nodes_compositor.py @@ -38,25 +38,62 @@ def expand_batch_frames(tensors: list[torch.Tensor]) -> list[torch.Tensor]: return frames -def stack_images(tensors: list[torch.Tensor]) -> torch.Tensor: - canvas = tensors[0][:1, :, :, :3].clone() - h, w = canvas.shape[1], canvas.shape[2] - for layer in tensors[1:]: - layer = layer[:1, :, :, :3] - lh = min(layer.shape[1], h) - lw = min(layer.shape[2], w) - canvas[:, :lh, :lw, :] = layer[:, :lh, :lw, :] - return canvas +def frame_alpha( + tensor: torch.Tensor, mask: torch.Tensor | None +) -> torch.Tensor | None: + alpha = tensor[:1, :, :, 3] if tensor.shape[-1] == 4 else None + if mask is None: + return alpha + h, w = tensor.shape[1], tensor.shape[2] + m = mask[:1].to(dtype=torch.float32) + if m.shape[1] != h or m.shape[2] != w: + m = torch.nn.functional.interpolate( + m.unsqueeze(1), size=(h, w), mode="bilinear" + ).squeeze(1) + inv = torch.clamp(1.0 - m, 0.0, 1.0) + return inv if alpha is None else alpha * inv -def input_fingerprints(tensors: list[torch.Tensor]) -> list[str]: +def frame_alphas( + tensors: list[torch.Tensor], masks: list[torch.Tensor] +) -> list[torch.Tensor | None]: + return [ + frame_alpha(tensor, masks[index] if index < len(masks) else None) + for index, tensor in enumerate(tensors) + ] + + +def layer_preview_tensor( + tensor: torch.Tensor, alpha: torch.Tensor | None +) -> torch.Tensor: + rgb = tensor[:1, :, :, :3] + if alpha is None: + return rgb + return torch.cat([rgb, alpha.unsqueeze(-1)], dim=-1) + + +def canvas_size(tensors: list[torch.Tensor]) -> tuple[int, int]: + return ( + max(tensor.shape[2] for tensor in tensors), + max(tensor.shape[1] for tensor in tensors), + ) + + +def input_fingerprints( + tensors: list[torch.Tensor], alphas: list[torch.Tensor | None] +) -> list[str]: fingerprints = [] - for tensor in tensors: - frame = tensor[0].detach().cpu().numpy() + for tensor, alpha in zip(tensors, alphas): + frame = tensor[0, :, :, :3].detach().cpu().numpy() frame8 = np.clip(np.rint(frame * 255.0), 0, 255).astype(np.uint8) digest = hashlib.sha256() digest.update(repr(tuple(tensor.shape)).encode()) digest.update(frame8.tobytes()) + if alpha is not None: + alpha8 = np.clip( + np.rint(alpha[0].detach().cpu().numpy() * 255.0), 0, 255 + ).astype(np.uint8) + digest.update(alpha8.tobytes()) fingerprints.append(digest.hexdigest()[:16]) return fingerprints @@ -147,7 +184,7 @@ def state_from_bboxes(tensors: list[torch.Tensor], slots: list) -> dict: } }) return { - "canvas": (tensors[0].shape[2], tensors[0].shape[1]), + "canvas": canvas_size(tensors), "layers": layers, "inputs": None, "background": {"color": "#ffffff", "opacity": 1.0, "visible": True}, @@ -179,6 +216,17 @@ def _parse_background(entry) -> dict | None: } +def _parse_order(value) -> list[int] | None: + if not isinstance(value, list) or not value: + return None + if not all( + isinstance(item, int) and not isinstance(item, bool) and item >= 0 + for item in value + ): + return None + return value + + def layer_state_provided(raw) -> bool: if isinstance(raw, dict): return bool(raw) @@ -219,6 +267,7 @@ def parse_layer_state(raw) -> dict | None: "layers": layers, "inputs": inputs, "background": _parse_background(state.get("background")), + "order": _parse_order(state.get("order")), } @@ -248,10 +297,18 @@ def _layer_params(entry, natural_w: int, natural_h: int) -> dict: } -def _prepare_layer_bitmap(tensor: torch.Tensor, params: dict) -> Image.Image: +def _prepare_layer_bitmap( + tensor: torch.Tensor, params: dict, alpha: torch.Tensor | None +) -> Image.Image: frame = tensor[0, :, :, :3].detach().cpu().numpy() rgb8 = np.clip(np.rint(frame * 255.0), 0, 255).astype(np.uint8) - img = Image.fromarray(rgb8, "RGB").convert("RGBA") + if alpha is None: + img = Image.fromarray(rgb8, "RGB").convert("RGBA") + else: + alpha8 = np.clip( + np.rint(alpha[0].detach().cpu().numpy() * 255.0), 0, 255 + ).astype(np.uint8) + img = Image.fromarray(np.dstack([rgb8, alpha8]), "RGBA") if params["flip_h"]: img = img.transpose(Image.Transpose.FLIP_LEFT_RIGHT) if params["flip_v"]: @@ -294,19 +351,29 @@ def _fill_background(canvas: np.ndarray, background: dict) -> np.ndarray: ) -def composite_from_state(tensors: list[torch.Tensor], state: dict) -> torch.Tensor: +def composite_from_state( + tensors: list[torch.Tensor], + state: dict, + alphas: list[torch.Tensor | None], +) -> torch.Tensor: cw, ch = state["canvas"] canvas = np.zeros((ch, cw, 4), dtype=np.float32) background = state.get("background") if background is not None and background["visible"] and background["opacity"] > 0: canvas = _fill_background(canvas, background) layers = state["layers"] - for index, tensor in enumerate(tensors): + order = state.get("order") or range(len(tensors)) + for index in order: + if index < 0 or index >= len(tensors): + continue + tensor = tensors[index] entry = layers[index] if index < len(layers) else None params = _layer_params(entry, tensor.shape[2], tensor.shape[1]) if not params["visible"]: continue - img = _prepare_layer_bitmap(tensor, params) + img = _prepare_layer_bitmap( + tensor, params, alphas[index] if index < len(alphas) else None + ) bx, by, bw, bh = placed_bounds( params["x"], params["y"], params["w"], params["h"], params["rotation"] ) @@ -321,8 +388,21 @@ def composite_from_state(tensors: list[torch.Tensor], state: dict) -> torch.Tens mode, canvas[y0:y1, x0:x1], region, params["opacity"] ) rgb = linear_to_srgb(np.clip(canvas[..., :3], 0.0, 1.0)) - rgb = rgb * np.clip(canvas[..., 3:4], 0.0, 1.0) - return torch.from_numpy(rgb.astype(np.float32)).unsqueeze(0) + alpha = np.clip(canvas[..., 3:4], 0.0, 1.0) + rgba = np.concatenate([rgb, alpha], axis=-1) + return torch.from_numpy(rgba.astype(np.float32)).unsqueeze(0) + + +OPAQUE_EPSILON = 1e-3 + + +def composite_outputs(out: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + if out.shape[-1] != 4: + return out, torch.zeros(out.shape[:3], dtype=torch.float32) + alpha = out[..., 3] + if bool((alpha >= 1.0 - OPAQUE_EPSILON).all()): + return out[..., :3], torch.zeros_like(alpha) + return out, torch.clamp(1.0 - alpha, 0.0, 1.0) class ImageCompositor(io.ComfyNode): @@ -330,7 +410,7 @@ class ImageCompositor(io.ComfyNode): def define_schema(cls): return io.Schema( node_id="ImageCompositor", - display_name="Image Compositor", + display_name="Create Layered Image", category="image", is_output_node=True, has_intermediate_output=True, @@ -343,50 +423,67 @@ class ImageCompositor(io.ComfyNode): min=1, max=50, ), - tooltip="Layers to composite. The first input is the bottom layer; each subsequent input is stacked above the previous one.", + tooltip="Layers to composite. The first image is the back layer; each subsequent image is stacked above the previous one.", + ), + io.Autogrow.Input( + "masks", + template=io.Autogrow.TemplatePrefix( + io.Mask.Input("mask"), + prefix="mask_", + min=0, + max=50, + ), + optional=True, + tooltip="Optional per-layer transparency masks, paired with image frames by index (mask_0 applies to the first frame). Masked areas (value 1) become transparent, multiplying with any alpha channel the image already carries.", ), io.MultiType.Input( "bboxes", [io.BoundingBox, io.Array, io.String], optional=True, - tooltip="Optional initial layout: bounding boxes, elements, or a JSON string, index-aligned with the image inputs (bboxes[0] places image_0). Inputs without a box keep their natural size at the origin. A saved compositor recipe that matches the current inputs takes priority.", + tooltip="Optional bounding boxes to initialize the layout, index-aligned with the image inputs (bboxes[0] places image_0). Images without a bounding box keep their natural size at the origin. A saved composition that matches the current set of inputs takes priority.", ), io.Compositor.Input( "compositor", - tooltip="Layer recipe saved by the compositor editor, replayed over the current inputs", + tooltip="Layered composition saved by the compositor editor.", ), ], outputs=[ - io.Image.Output(), + io.Image.Output( + tooltip="Composited image. Carries an alpha channel when the composite has transparent areas (e.g. hidden background), otherwise plain RGB." + ), + io.Mask.Output( + tooltip="Transparency of the composite (1 = fully transparent). All zeros when the composite is opaque." + ), ], ) @classmethod - def execute(cls, images: io.Autogrow.Type = None, compositor: io.Compositor.Type = None, bboxes: io.MultiType.Type = None) -> io.NodeOutput: + def execute(cls, images: io.Autogrow.Type = None, masks: io.Autogrow.Type = None, compositor: io.Compositor.Type = None, bboxes: io.MultiType.Type = None) -> io.NodeOutput: tensors = expand_batch_frames(sort_autogrow_images(images)) + mask_frames = expand_batch_frames(sort_autogrow_images(masks)) + alphas = frame_alphas(tensors, mask_frames) layer_refs = [] - for tensor in tensors: - layer_refs.extend(UI.PreviewImage(tensor, cls=cls).values) + for tensor, alpha in zip(tensors, alphas): + layer_refs.extend( + UI.PreviewImage(layer_preview_tensor(tensor, alpha), cls=cls).values + ) - fp = input_fingerprints(tensors) + fp = input_fingerprints(tensors, alphas) raw_state = compositor state = parse_layer_state(raw_state) replay = bool(state is not None and tensors and state["inputs"] == fp) - slots = ( - layout_bboxes(bboxes, tensors[0].shape[2], tensors[0].shape[1]) - if tensors - else [] - ) + slots = layout_bboxes(bboxes, *canvas_size(tensors)) if tensors else [] if replay: - out = composite_from_state(tensors, state) - elif tensors and any(slot is not None for slot in slots): - out = composite_from_state(tensors, state_from_bboxes(tensors, slots)) + out = composite_from_state(tensors, state, alphas) elif tensors: - out = stack_images(tensors) + out = composite_from_state( + tensors, state_from_bboxes(tensors, slots), alphas + ) else: out = torch.zeros((1, 64, 64, 3), dtype=torch.float32) state_stale = layer_state_provided(raw_state) and not replay + out, mask = composite_outputs(out) ui_dict = UI.PreviewImage(out, cls=cls).as_dict() ui_dict["compositor_layers"] = layer_refs @@ -394,7 +491,7 @@ class ImageCompositor(io.ComfyNode): ui_dict["compositor_bboxes"] = bbox_ui_entries(slots, len(tensors)) if state_stale: ui_dict["compositor_state_stale"] = [True] - return io.NodeOutput(out, ui=ui_dict) + return io.NodeOutput(out, mask, ui=ui_dict) class CompositorExtension(ComfyExtension):