diff --git a/ai_diffusion/backend/resources.py b/ai_diffusion/backend/resources.py index 2830ee27d..ea77fc60a 100644 --- a/ai_diffusion/backend/resources.py +++ b/ai_diffusion/backend/resources.py @@ -181,7 +181,7 @@ def has_controlnet_inpaint(self): @property def supports_regions(self): - return self in [Arch.sd15, Arch.sdxl, Arch.illu, Arch.illu_v] + return self in [Arch.sd15, Arch.sdxl, Arch.illu, Arch.illu_v, Arch.anima] @property def supports_lcm(self): @@ -738,6 +738,7 @@ def is_required(kind: ResourceKind, arch: Arch, identifier: ControlMode | Upscal resource_id(ResourceKind.controlnet, Arch.illu, ControlMode.universal): ["union-sdxl", "xinsirunion"], resource_id(ResourceKind.controlnet, Arch.illu_v, ControlMode.universal): ["union-sdxl", "xinsirunion"], resource_id(ResourceKind.controlnet, Arch.anima, ControlMode.universal): ["anima*lllite*any"], + resource_id(ResourceKind.controlnet, Arch.anima, ControlMode.segmentation): ["anima-lllite-region-cn", "anima*lllite*region"], resource_id(ResourceKind.controlnet, Arch.flux, ControlMode.universal): ["flux.1-dev-controlnet-union-pro-2.0", "flux.1-dev-controlnet-union-pro", "flux.1-dev-controlnet-union", "flux1devcontrolnetunion"], resource_id(ResourceKind.controlnet, Arch.qwen, ControlMode.universal): ["qwen-image-instantx-controlnet-union"], resource_id(ResourceKind.controlnet, Arch.sd15, ControlMode.scribble): ["control_v11p_sd15_scribble", "control_lora_rank128_v11p_sd15_scribble"], diff --git a/ai_diffusion/model/control.py b/ai_diffusion/model/control.py index 961be2cd9..c3a67f938 100644 --- a/ai_diffusion/model/control.py +++ b/ai_diffusion/model/control.py @@ -2,15 +2,16 @@ import json from pathlib import Path -from typing import Any, NamedTuple +from typing import Any, ClassVar, NamedTuple from PyQt5.QtCore import QObject, Qt, QUuid, pyqtSignal +from PyQt5.QtGui import QColor from .. import util from ..backend import resources from ..backend.api import ControlInput from ..backend.resources import Arch, ControlMode, ResourceKind, resource_id -from ..image import Bounds, Extent, Image +from ..image import BlendMode, Bounds, Extent, Image from ..layer import Layer, LayerType from ..localization import translate as _ from ..util import PluginError @@ -23,6 +24,20 @@ class ControlLayer(QObject, ObservableProperties): max_preset_value = 4 strength_multiplier = 50 clip_vision_extent = Extent(224, 224) + segmentation_colors: ClassVar[list[tuple[int, int, int]]] = [ + (120, 120, 120), + (180, 120, 120), + (120, 180, 120), + (120, 120, 180), + (180, 180, 120), + (180, 120, 180), + (120, 180, 180), + (220, 140, 100), + (140, 220, 100), + (100, 140, 220), + (220, 100, 140), + (100, 220, 140), + ] mode = Property(ControlMode.reference, persist=True, setter="set_mode") layer_id = Property(QUuid(), persist=True) @@ -125,6 +140,9 @@ def to_api(self, bounds: Bounds | None = None, time: int | None = None): image = layer.get_pixels(bounds, time) + if self.mode is ControlMode.segmentation: + image.make_opaque(background=Qt.GlobalColor.white) + if self.mode.is_lines or self.mode is ControlMode.stencil: image.make_opaque(background=Qt.GlobalColor.white) @@ -139,6 +157,56 @@ def to_api(self, bounds: Bounds | None = None, time: int | None = None): strength = self.strength / self.strength_multiplier return ControlInput(self.mode, image, strength, (self.start, self.end)) + def generate_segmentation(self): + assert self.mode is ControlMode.segmentation + + ok, msg = self._model.document.check_color_mode() + if not ok and msg: + self._model.report_error(msg) + return + + try: + bounds = Bounds.from_extent(self._model.document.extent) + image = self._segmentation_image_from_regions(bounds) + if image is None: + self._model.report_error(_("Text prompt regions have not been set up.")) + return + + layer = self._model.layers.create(f"[Control] {self.mode.text}", image, bounds) + self.layer_id = layer.id + except Exception as e: + self._model.report_error(util.log_error(e)) + else: + self._model.clear_error() + + def _segmentation_image_from_regions(self, bounds: Bounds): + from .region import RegionLink + + image = Image.create(bounds.extent, fill=Qt.GlobalColor.white) + has_region_layer = False + root = self._model.active_regions + + layers = [ + layer + for layer in root.layers.all + if root.find_linked(layer, RegionLink.direct) is not None + and Bounds.intersection(bounds, layer.compute_bounds()).area > 0 + ] + + for index, layer in enumerate(layers): + color = self.segmentation_colors[index % len(self.segmentation_colors)] + region_image = self._segmentation_region_image(layer, bounds, color) + image.draw_image(region_image, blend=BlendMode.alpha) + has_region_layer = True + + return image if has_region_layer else None + + def _segmentation_region_image(self, layer: Layer, bounds: Bounds, color: tuple[int, int, int]): + mask = layer.get_mask(bounds) + image = Image.create(bounds.extent, fill=QColor(*color, 255)) + image._qimage.setAlphaChannel(mask._qimage) + return image + def generate(self): self._generate_job = self._model.generate_control_layer(self) self.has_active_job = True diff --git a/ai_diffusion/presets/control.json b/ai_diffusion/presets/control.json index d4561ff1e..99d5a5a81 100644 --- a/ai_diffusion/presets/control.json +++ b/ai_diffusion/presets/control.json @@ -199,5 +199,41 @@ "end": 1.0 } ] + }, + "segmentation": { + "all": [ + { + "strength": 0.7, + "start": 0.0, + "end": 0.5 + }, + { + "strength": 1.0, + "start": 0.0, + "end": 0.8 + }, + { + "strength": 1.0, + "start": 0.0, + "end": 1.0 + } + ], + "anima": [ + { + "strength": 0.75, + "start": 0.0, + "end": 0.45 + }, + { + "strength": 1.0, + "start": 0.0, + "end": 0.45 + }, + { + "strength": 1.5, + "start": 0.0, + "end": 0.45 + } + ] } -} \ No newline at end of file +} diff --git a/ai_diffusion/ui/control.py b/ai_diffusion/ui/control.py index 7981714e5..07fccc181 100644 --- a/ai_diffusion/ui/control.py +++ b/ai_diffusion/ui/control.py @@ -73,6 +73,11 @@ def __init__( ) self.generate_tool_button.clicked.connect(control.generate) + self.generate_regions_tool_button = _create_generate_regions_button( + self, Qt.ToolButtonStyle.ToolButtonIconOnly + ) + self.generate_regions_tool_button.clicked.connect(control.generate_segmentation) + self.add_pose_tool_button = _create_add_pose_button( self, Qt.ToolButtonStyle.ToolButtonIconOnly ) @@ -91,6 +96,7 @@ def __init__( bar_layout.addWidget(self.mode_select) bar_layout.addWidget(self.layer_select, 3) bar_layout.addWidget(self.generate_tool_button) + bar_layout.addWidget(self.generate_regions_tool_button) bar_layout.addWidget(self.add_pose_tool_button) bar_layout.addWidget(self.preset_slider, 1) bar_layout.addWidget(self.error_text, 3) @@ -129,6 +135,11 @@ def __init__( ) self.generate_button.clicked.connect(control.generate) + self.generate_regions_button = _create_generate_regions_button( + self.extended_widget, Qt.ToolButtonStyle.ToolButtonTextBesideIcon + ) + self.generate_regions_button.clicked.connect(control.generate_segmentation) + self.add_pose_button = _create_add_pose_button( self.extended_widget, Qt.ToolButtonStyle.ToolButtonTextBesideIcon ) @@ -141,6 +152,7 @@ def __init__( actions_layout = QHBoxLayout() actions_layout.addWidget(self.custom_checkbox, stretch=1) actions_layout.addWidget(self.generate_button) + actions_layout.addWidget(self.generate_regions_button) actions_layout.addWidget(self.add_pose_button) extended_layout.addLayout(actions_layout) @@ -228,6 +240,7 @@ def _add_pose_character(self): def _update_visibility(self): is_small = self.width() < 420 is_pose = self._control.mode is ControlMode.pose + is_segmentation = self._control.mode is ControlMode.segmentation is_edit = root.active_model.arch.supports_edit def controls(): @@ -236,6 +249,12 @@ def controls(): self.expand_button.setVisible(self._control.is_supported and not is_edit) self.generate_button.setVisible(self._control.can_generate and is_small) self.generate_tool_button.setVisible(self._control.can_generate and not is_small) + self.generate_regions_button.setVisible( + self._control.is_supported and is_segmentation and is_small + ) + self.generate_regions_tool_button.setVisible( + self._control.is_supported and is_segmentation and not is_small + ) self.add_pose_button.setVisible(is_pose and is_small) self.add_pose_tool_button.setVisible(is_pose and not is_small) self.range_label.setVisible(self._control.has_range) @@ -272,6 +291,8 @@ def _update_strength(self): def _update_job_active(self): self.generate_button.setEnabled(not self._control.has_active_job) self.generate_tool_button.setEnabled(not self._control.has_active_job) + self.generate_regions_button.setEnabled(not self._control.has_active_job) + self.generate_regions_tool_button.setEnabled(not self._control.has_active_job) self.layer_select.setEnabled(not self._control.has_active_job) def _update_custom_values(self): @@ -309,6 +330,15 @@ def _create_generate_button(parent, style: Qt.ToolButtonStyle): return button +def _create_generate_regions_button(parent, style: Qt.ToolButtonStyle): + button = QToolButton(parent) + button.setToolButtonStyle(style) + button.setText(_("From Regions")) + button.setIcon(theme.icon("region-prompt")) + button.setToolTip(_("Generate segmentation control layer from current regions")) + return button + + def _create_add_pose_button(parent, style: Qt.ToolButtonStyle): button = QToolButton(parent) button.setToolButtonStyle(style)