diff --git a/ArtKrit/.gitignore b/ArtKrit/.gitignore new file mode 100644 index 0000000..83d5583 --- /dev/null +++ b/ArtKrit/.gitignore @@ -0,0 +1,177 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# UV +# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +#uv.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/latest/usage/project/#working-with-version-control +.pdm.toml +.pdm-python +.pdm-build/ + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +#.idea/ + +# PyPI configuration file +.pypirc + +res/* +detection_res.png +temp/* +*.DS_Store/* +*.DS_Store \ No newline at end of file diff --git a/ArtKrit/README.md b/ArtKrit/README.md new file mode 100644 index 0000000..886aab8 --- /dev/null +++ b/ArtKrit/README.md @@ -0,0 +1,92 @@ +# ArtKrit +ArtKrit is a plugin for Krita that helps artists enhance their drawing skills by scaffolding the process of replicating a reference image into three steps: composition, value, and color. At each stage, ArtKrit generates adaptive composition lines and provides feedback on value and color accuracy to help artists refine their work. + +fig_teaser + +**Bottom row**: Computational guidance and feedback provided by our system at each step. We offer object-based composition lines to assist with spatial positioning, and we visualize differences in value and color with verbal suggestions to guide the user. + +*Reference image: "Interior Practice // Kitchen" by Loish (2021, digital).* + +## Installation + + +### Krita +1. Install Krita from the official website: https://krita.org/en/download/ (tested on version 5.2.9) +2. To facilitate debugging, you can add the path to your Krita binary to your bash or zsh profile. On Mac, it should look like this (it would be bash for older MacOS versions): + ```bash + echo 'export PATH="/Applications/krita.app/Contents/MacOS/:$PATH"' >> ~/.zshrc + source ~/.zshrc + ``` + This allows you to run Krita from the terminal by typing `krita`. All the output from Krita and the plugin will be printed to the terminal. + + +### File Structure Setup +1. On Mac, the Python plugin folder is located at `~/Library/Application Support/Krita/pykrita/`. Navigate to this folder and `git clone` this repository. Note that for MacOS, ~/Library/Application Support and /Library/Application Support are different folders. If you don't find the Krita folder, make sure you are in the Application Support for your user. + +2. Under the `pykrita` folder, create a `artkrit.desktop` file. The file structure now should look like this: + ``` + pykrita/ + ArtKrit/ + ... + artkrit.desktop + ``` + +3. In the `artkrit.desktop` file, add the following content: + ```ini + # File: artkrit.desktop + [Desktop Entry] + Type=Service + ServiceTypes=Krita/PythonPlugin + X-KDE-Library=ArtKrit + X-Python-2-Compatible=false + X-Krita-Manual=Manual.html + Name=ArtKrit + Comment=Docker for ArtKrit + ``` + + +### Python Plugin Setup + +1. Pick your favorite virtual environment tool (e.g. `uv`, `venv`, `conda`, etc.) and create a new environment with `python==3.10` at your home directory (`~`). + - Make sure you use `python==3.10` for compatibility with Krita 5.2.9. + - Name your environment `ddraw` for consistency. If you name it something else or place it elsewhere, update the system path at the top of `artkrit.py` and `value_color.py`. + - I recommend using [`uv`](https://docs.astral.sh/uv/) to manage your environments for its simplicity and speed. + - To install uv (on MacOS), `curl -LsSf https://astral.sh/uv/install.sh | sh` + - To install the virtual environment, `uv venv ddraw --python 3.10` + +2. Activate your environment (e.g., `source ddraw/bin/activate`). Now navigate (`cd`) to `~/Library/Application\ Support/Krita/pykrita/ArtKrit` + +3. First, install pytorch with: + ```bash + pip install torch torchvision torchaudio + + # If you're using `uv`, you can use the following command: + uv pip install torch torchvision torchaudio + ``` + +4. Then install the other required packages: + ```bash + pip install -r requirements.txt + + # If you're using `uv`, you can use the following command: + uv pip install -r requirements.txt + ``` + +## Running the Plugin +1. In one terminal, run Krita by typing `krita` in the terminal. This will allow you to see the output from the plugin. Directly opening Krita also works, but you won't see the output. + +2. In another terminal, activate your environment and navigate to the `ArtKrit` folder. Then start the python server with: + ```bash + python script/composition/server.py + ``` + Note, if you get an error, try running it with the specific python version: `python3.10 server.py` + +3. On the first launch, enable the plugin by going to Preferences (cmd+`,`), scrolling down, selecting Python Plugin Manager, and checking the ArtKrit box. Then, relaunch Krita. The docker (window) for the plug-in can be found under Settings > Dockers > ArtKrit. +4. When setting up a Krita document, it is recommended to set it as the same size as the reference image. This will ensure that the plugin works as intended. +5. Make sure to click `Set Reference Image` button every time you reopen Krita. +6. If inferencing time for generating adaptive grids is too long, you can try to use smaller models listed in `run_models.py`. Note that smaller models will not be as performant. + + +## Helpful Resources +Krita API Documentation: [https://api.kde.org/krita/html/](https://api.kde.org/krita/html/) +Guide for plugins: [https://docs.krita.org/en/user_manual/python_scripting/krita_python_plugin_howto.html](https://docs.krita.org/en/user_manual/python_scripting/krita_python_plugin_howto.html) diff --git a/ArtKrit/__init__.py b/ArtKrit/__init__.py new file mode 100644 index 0000000..ea8613d --- /dev/null +++ b/ArtKrit/__init__.py @@ -0,0 +1 @@ +from .artkrit import * \ No newline at end of file diff --git a/ArtKrit/artkrit.py b/ArtKrit/artkrit.py new file mode 100644 index 0000000..23673a0 --- /dev/null +++ b/ArtKrit/artkrit.py @@ -0,0 +1,1078 @@ +from krita import DockWidget, DockWidgetFactory, DockWidgetFactoryBase, Krita, InfoObject +from PyQt5.QtCore import Qt, QPointF, QMimeData, QEventLoop, QTimer +from PyQt5.QtWidgets import ( + QWidget, QVBoxLayout, QPushButton, QLabel, QFileDialog, QLineEdit, QHBoxLayout, QSlider, QSpinBox, QSizePolicy, + QTabWidget, QScrollArea, QGroupBox, QDialog +) +from PyQt5.QtGui import QImage, QPixmap, QPainter, QPen, QColor, QPainterPath, QGuiApplication, QClipboard +from ArtKrit.script.value_color.value_color import ValueColor +import json +from datetime import datetime +from ArtKrit.script.composition.run_models import init_models, detect, segment +from ArtKrit.script.composition.composition_utils import process_image_direct, regenerate_lines_direct + +import os +import sys +sys.path.append(os.path.expanduser("~/ddraw/lib/python3.10/site-packages")) + +class PreviewDialog(QDialog): + """Popup dialog for showing the reference image with overlays""" + def __init__(self, parent, reference_image): + super().__init__(parent) + self.parent_widget = parent + self.reference_image = reference_image + self.setWindowTitle("Reference Image Preview") + self.resize(800, 600) + + layout = QVBoxLayout() + self.preview_label = QLabel() + self.preview_label.setAlignment(Qt.AlignCenter) + self.preview_label.setScaledContents(False) + layout.addWidget(self.preview_label) + + self.setLayout(layout) + self.update_preview() + + def update_preview(self): + """Update the preview with current overlays""" + if self.reference_image: + pixmap = QPixmap.fromImage(self.reference_image) + pixmap = self.parent_widget.draw_overlays_on_pixmap(pixmap) + + # Scale to fit dialog while maintaining aspect ratio + scaled_pixmap = pixmap.scaled( + self.preview_label.size(), + Qt.KeepAspectRatio, + Qt.SmoothTransformation + ) + self.preview_label.setPixmap(scaled_pixmap) + + def resizeEvent(self, event): + """Handle resize events to update preview""" + super().resizeEvent(event) + self.update_preview() + + +class ArtKrit(DockWidget): + def __init__(self): + super().__init__() + self.setWindowTitle("ArtKrit") + self.preview_image = None + self.image_file_path = None + self.compose_lines = [] + self.cached_points = [] # Store points for regeneration + self.cached_polygon_contours = [] # Store polygon contours + self.value_color = ValueColor(self) + + # Track overlay visibility states + self.thirds_visible = False + self.cross_visible = False + self.circle_visible = False + self.adaptive_grid_visible = False + self.contours_visible = False + + print("[Plugin] Initializing models...") + self.detector, self.segmentator, self.processor = init_models() + print("[Plugin] Models initialized successfully") + + # Reference to popup dialog + self.preview_dialog = None + + self.setUI() + + + def setUI(self): + # Main widget and layout + self.main_widget = QWidget() + self.setWidget(self.main_widget) + self.main_layout = QVBoxLayout() + self.main_layout.setAlignment(Qt.AlignTop) # Align widgets to the top + self.main_widget.setLayout(self.main_layout) + + # Create tab widget + self.tab_widget = QTabWidget() + self.main_layout.addWidget(self.tab_widget) + + # Create first tab (Composition Grid) + self.create_composition_tab() + self.value_color.create_value_tab() + self.value_color.create_color_tab() + + # Add tabs to tab widget + self.tab_widget.addTab(self.composition_tab, "Composition") + self.tab_widget.addTab(self.value_color.value_tab, "Value") + self.tab_widget.addTab(self.value_color.color_tab, "Color") + + # Create a scroll area and set the main widget + scroll_area = QScrollArea() + scroll_area.setWidgetResizable(True) + scroll_area.setWidget(self.main_widget) + self.setWidget(scroll_area) + + + def create_composition_tab(self): + self.composition_tab = QWidget() + self.composition_layout = QVBoxLayout() + self.composition_layout.setAlignment(Qt.AlignTop) + self.composition_tab.setLayout(self.composition_layout) + + # Add a button to set reference image + self.set_reference_image_btn = QPushButton("Set Reference Image") + self.set_reference_image_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) + self.set_reference_image_btn.clicked.connect(self.set_reference_image) + self.composition_layout.addWidget(self.set_reference_image_btn) + + # Add preview area for reference image + self.preview_group = QGroupBox("Reference Preview") + self.preview_layout = QVBoxLayout() + self.preview_group.setLayout(self.preview_layout) + + # Preview label for showing the image + self.preview_label = QLabel() + self.preview_label.setAlignment(Qt.AlignCenter) + self.preview_label.setMinimumHeight(200) + self.preview_label.setMaximumHeight(300) + self.preview_label.setScaledContents(False) + self.preview_label.setStyleSheet("QLabel { background-color: #2a2a2a; border: 1px solid #555; }") + self.preview_layout.addWidget(self.preview_label) + + # Pop out button + self.popout_btn = QPushButton("Pop Out Preview") + self.popout_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) + self.popout_btn.clicked.connect(self.toggle_preview_dialog) + self.popout_btn.setEnabled(False) + self.preview_layout.addWidget(self.popout_btn) + + self.composition_layout.addWidget(self.preview_group) + + # Create a group box for predefined grids + self.predefined_grids_group = QGroupBox("Predefined Grid") + self.predefined_grids_layout = QVBoxLayout() + self.predefined_grids_group.setLayout(self.predefined_grids_layout) + + # Add Rule of Thirds button + self.thirds_btn = QPushButton("Toggle Rule of Thirds Grid") + self.thirds_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) + self.thirds_btn.clicked.connect(self.toggle_canvas_thirds) + self.predefined_grids_layout.addWidget(self.thirds_btn) + + # Add Cross Grid button + self.cross_btn = QPushButton("Toggle Cross Grid") + self.cross_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) + self.cross_btn.clicked.connect(self.toggle_canvas_cross) + self.predefined_grids_layout.addWidget(self.cross_btn) + + # Add Circle Grid button + self.circle_btn = QPushButton("Toggle Circle Grid") + self.circle_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) + self.circle_btn.clicked.connect(self.toggle_canvas_circle) + self.predefined_grids_layout.addWidget(self.circle_btn) + + # Add the group box to the composition layout + self.composition_layout.addWidget(self.predefined_grids_group) + + # Create a group box for adaptive grid settings + self.adaptive_grid_group = QGroupBox("Adaptive Grid") + self.adaptive_grid_layout = QVBoxLayout() + self.adaptive_grid_group.setLayout(self.adaptive_grid_layout) + + ## add a text input field for the user to input text prompt for GroundingDINO + self.text_prompt_widget = QWidget() # Create a widget to hold the label and input + self.text_prompt_widget.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) # Prevent vertical stretching + self.text_prompt_layout = QHBoxLayout() # Create a horizontal layout + self.text_prompt_layout.setContentsMargins(0, 3, 0, 0) # Small top margin for spacing + self.text_prompt_label = QLabel("Text Prompt") # Create a label for the text input + self.text_prompt_input = QLineEdit() # Create a text input field + self.text_prompt_layout.addWidget(self.text_prompt_label) # Add the label to the horizontal layout + self.text_prompt_layout.addWidget(self.text_prompt_input) # Add the input field to the horizontal layout + self.text_prompt_widget.setLayout(self.text_prompt_layout) # Set the layout for the widget + self.adaptive_grid_layout.addWidget(self.text_prompt_widget) # Add the widget to the adaptive grid layout + + ## add a numeric slider with value between 0 and 20, label it with "Polygon Epsilon". make them in the same row + self.polygon_epsilon_widget = QWidget() + self.polygon_epsilon_widget.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) # Prevent vertical stretching + self.polygon_epsilon_layout = QHBoxLayout() # Change to QHBoxLayout to place slider and spinbox in the same row + self.polygon_epsilon_layout.setContentsMargins(0, 3, 0, 0) # Small top margin for spacing + self.polygon_epsilon_label = QLabel("Polygon Epsilon") # Initial label with default value + self.polygon_epsilon_slider = QSlider(Qt.Horizontal) + self.polygon_epsilon_slider.setValue(8) # Set a default value + self.polygon_epsilon_slider.setMinimum(0) + self.polygon_epsilon_slider.setMaximum(20) + + self.polygon_epsilon_spinbox = QSpinBox() # Create a spinbox + self.polygon_epsilon_spinbox.setMinimum(0) + self.polygon_epsilon_spinbox.setMaximum(20) + self.polygon_epsilon_spinbox.setValue(8) # Set a default value + + # Connect slider and spinbox to update each other + self.polygon_epsilon_slider.valueChanged.connect(self.polygon_epsilon_spinbox.setValue) + self.polygon_epsilon_spinbox.valueChanged.connect(self.polygon_epsilon_slider.setValue) + + self.polygon_epsilon_layout.addWidget(self.polygon_epsilon_label) + self.polygon_epsilon_layout.addWidget(self.polygon_epsilon_slider) + self.polygon_epsilon_layout.addWidget(self.polygon_epsilon_spinbox) + self.polygon_epsilon_widget.setLayout(self.polygon_epsilon_layout) + self.adaptive_grid_layout.addWidget(self.polygon_epsilon_widget) + + ## add a slider to change the number of lines shown in the composition grid + self.grid_lines_widget = QWidget() + self.grid_lines_widget.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) # Prevent vertical stretching + self.grid_lines_layout = QHBoxLayout() # Change to QHBoxLayout to place slider and spinbox in the same row + self.grid_lines_layout.setContentsMargins(0, 3, 0, 3) # Small margins for spacing + self.grid_lines_label = QLabel("Number of Grid Lines") # Initial label without default value + self.grid_lines_slider = QSlider(Qt.Horizontal) + self.grid_lines_slider.setValue(2) # Set a default value + self.grid_lines_slider.setMinimum(1) # Minimum 1 line + self.grid_lines_slider.setMaximum(10) # Maximum 10 lines + + self.grid_lines_spinbox = QSpinBox() # Create a spinbox + self.grid_lines_spinbox.setMinimum(1) + self.grid_lines_spinbox.setMaximum(10) + self.grid_lines_spinbox.setValue(2) # Set a default value + + # Connect slider and spinbox to update each other + self.grid_lines_slider.valueChanged.connect(self.grid_lines_spinbox.setValue) + self.grid_lines_spinbox.valueChanged.connect(self.grid_lines_slider.setValue) + + self.grid_lines_slider.valueChanged.connect(self.draw_composition_lines) + self.grid_lines_slider.valueChanged.connect(self.update_preview) + self.grid_lines_layout.addWidget(self.grid_lines_label) + self.grid_lines_layout.addWidget(self.grid_lines_slider) + self.grid_lines_layout.addWidget(self.grid_lines_spinbox) + self.grid_lines_widget.setLayout(self.grid_lines_layout) + self.adaptive_grid_layout.addWidget(self.grid_lines_widget) + + # Create composition grid button + self.canvas_circle_btn = QPushButton("Generate Adaptive Grid") + self.canvas_circle_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) # Prevent vertical stretching + self.canvas_circle_btn.clicked.connect(self.draw_grid) + self.adaptive_grid_layout.addWidget(self.canvas_circle_btn) + + # Add NEW button to regenerate lines from current points + self.regenerate_lines_btn = QPushButton("Regenerate Lines from Points") + self.regenerate_lines_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) + self.regenerate_lines_btn.clicked.connect(self.regenerate_lines_from_points) + self.adaptive_grid_layout.addWidget(self.regenerate_lines_btn) + + # Add a button that to toggle the visibility of the adaptive grid + self.toggle_adaptive_grid_btn = QPushButton("Toggle Adaptive Grid") + self.toggle_adaptive_grid_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) # Prevent vertical stretching + self.toggle_adaptive_grid_btn.clicked.connect(self.toggle_adaptive_grid) + self.adaptive_grid_layout.addWidget(self.toggle_adaptive_grid_btn) + + # Add the adaptive grid group to the main composition layout + self.composition_layout.addWidget(self.adaptive_grid_group) + + # Add a button to show the contours + self.show_contours_btn = QPushButton("Get Composition Feedback") + self.show_contours_btn.setSizePolicy(QSizePolicy.Preferred, QSizePolicy.Fixed) # Prevent vertical stretching + self.show_contours_btn.clicked.connect(self.toggle_contours) + self.composition_layout.addWidget(self.show_contours_btn) + + def process_image(self, image_path, text_prompt, custom_rectangles, polygon_epsilon): + """ + Process image using direct model calls (no server needed). + + This replaces the old HTTP POST to /process_image endpoint. + """ + try: + from PIL import Image + from .script.composition.composition_utils import load_image, DetectionResult + import time + + t0 = time.time() + + # Load image + if isinstance(image_path, str): + image = load_image(image_path) + else: + image = image_path # Already a PIL Image + + # Parse labels + labels = [l.strip() for l in text_prompt.split(",") if l.strip()] + threshold_bbox = 0.3 + polygon_refinement = True + + if (not labels) and (len(custom_rectangles) == 0): + return {"error": "No labels or custom rectangles provided"} + + # Detect + t_load = time.time() + print(f"[Direct] Calling GroundingDINO (labels={labels}, threshold={threshold_bbox})") + + # Import detect and segment HERE to avoid circular imports + from .script.composition.run_models import detect, segment + + detections = detect(image, labels, self.detector, threshold_bbox) + t_detect = time.time() + print(f"[Direct] GroundingDINO finished in {t_detect - t_load:.2f}s") + + # Add custom rectangles + for custom_rectangle in custom_rectangles: + detections.append(DetectionResult.from_dict({ + "score": 1.0, + "label": "custom_rectangle", + "box": { + "xmin": int(custom_rectangle[0]), + "ymin": int(custom_rectangle[1]), + "xmax": int(custom_rectangle[2]), + "ymax": int(custom_rectangle[3]) + } + })) + + # Segment + print("[Direct] Calling SAM model for segmentation") + detections = segment(image, detections, self.segmentator, self.processor, None, polygon_refinement) + t_segment = time.time() + print(f"[Direct] SAM finished in {t_segment - t_detect:.2f}s") + + # Now call the simplified process function + from .script.composition.composition_utils import process_image_direct + result = process_image_direct(image, detections, polygon_epsilon) + + if "error" in result: + print(f"[Plugin] Error: {result['error']}") + return None + + return result + + except Exception as e: + print(f"[Plugin] Error processing image: {e}") + import traceback + traceback.print_exc() + return None + + + def regenerate_lines(self, points, polygon_contours): + """ + Regenerate composition lines from manually adjusted points. + + This replaces the old HTTP POST to /regenerate_lines endpoint. + """ + try: + lines_list = regenerate_lines_direct(points, polygon_contours) + + return { + 'composition_lines': lines_list, + 'num_points': len(points), + 'num_lines': len(lines_list) + } + + except Exception as e: + print(f"[Plugin] Error regenerating lines: {e}") + import traceback + traceback.print_exc() + return None + + def draw_overlays_on_pixmap(self, pixmap): + """Draw all active overlays on a pixmap""" + if pixmap.isNull(): + return pixmap + + # Create a copy to draw on + result = pixmap.copy() + painter = QPainter(result) + pen = QPen(QColor(0, 255, 0)) + pen.setWidth(max(3, result.width() // 200)) # Scale pen width + painter.setPen(pen) + + width = result.width() + height = result.height() + + # Draw thirds grid + if self.thirds_visible: + for i in range(1, 3): + x = width * i / 3 + painter.drawLine(int(x), 0, int(x), height) + for i in range(1, 3): + y = height * i / 3 + painter.drawLine(0, int(y), width, int(y)) + + # Draw cross grid + if self.cross_visible: + x = width / 2 + painter.drawLine(int(x), 0, int(x), height) + y = height / 2 + painter.drawLine(0, int(y), width, int(y)) + + # Draw circle grid + if self.circle_visible: + center_x = width / 2 + center_y = height / 2 + radius = min(width, height) / 4 + painter.drawEllipse(int(center_x - radius), int(center_y - radius), + int(radius * 2), int(radius * 2)) + + # Draw contours + if self.contours_visible and self.cached_polygon_contours: + pen.setColor(QColor(0, 0, 255)) + painter.setPen(pen) + document = Krita.instance().activeDocument() + if document: + scale_x = width / document.width() + scale_y = height / document.height() + for polygon in self.cached_polygon_contours: + path = QPainterPath() + if polygon: + first_point = polygon[0] + path.moveTo(first_point[0] * scale_x, first_point[1] * scale_y) + for point in polygon[1:]: + path.lineTo(point[0] * scale_x, point[1] * scale_y) + path.closeSubpath() + painter.drawPath(path) + + # Draw adaptive grid lines + if self.adaptive_grid_visible and self.compose_lines: + pen.setColor(QColor(0, 255, 0)) + painter.setPen(pen) + document = Krita.instance().activeDocument() + if document: + scale_x = width / document.width() + scale_y = height / document.height() + num_lines_to_draw = min(self.grid_lines_slider.value(), len(self.compose_lines)) + for line in self.compose_lines[:num_lines_to_draw]: + p1, p2 = line + painter.drawLine(int(p1[0] * scale_x), int(p1[1] * scale_y), + int(p2[0] * scale_x), int(p2[1] * scale_y)) + + painter.end() + self.value_color.export_pixmap(result, "overlayed composition preview") + return result + + + def update_preview(self): + """Update the preview in both the dock and popup dialog""" + if self.preview_image: + pixmap = QPixmap.fromImage(self.preview_image) + pixmap = self.draw_overlays_on_pixmap(pixmap) + + # Update dock preview + scaled_pixmap = pixmap.scaled( + self.preview_label.width(), + self.preview_label.height(), + Qt.KeepAspectRatio, + Qt.SmoothTransformation + ) + self.preview_label.setPixmap(scaled_pixmap) + + # Update popup dialog if it exists + if self.preview_dialog and self.preview_dialog.isVisible(): + self.preview_dialog.update_preview() + + + def toggle_preview_dialog(self): + """Toggle the preview popup dialog""" + if self.preview_dialog is None or not self.preview_dialog.isVisible(): + self.preview_dialog = PreviewDialog(self, self.preview_image) + self.preview_dialog.show() + self.popout_btn.setText("Close Pop Out") + self.value_color.append_log_entry("preview popup open", "Opened preview popup dialog") + else: + self.preview_dialog.close() + self.preview_dialog = None + self.popout_btn.setText("Pop Out Preview") + self.value_color.append_log_entry("preview popup close", "Closed preview popup dialog") + + + def read_points_from_layer(self): + """Read point positions from the Points vector layer""" + document = Krita.instance().activeDocument() + if not document: + return [] + + points_layer = document.nodeByName('Points') + if not points_layer or points_layer.type() != "vectorlayer": + print("Points layer not found or not a vector layer") + return [] + + points = [] + points_per_inch = 72.0 + + # Extract circle centers from SVG shapes + for shape in points_layer.shapes(): + if shape.type() == "KoPathShape": + # Get the bounding box of the circle + bbox = shape.boundingBox() + # Calculate center point + center_x = (bbox.topLeft().x() + bbox.bottomRight().x()) / 2 + center_y = (bbox.topLeft().y() + bbox.bottomRight().y()) / 2 + # Convert from points to pixels + x = int(center_x * document.xRes() / points_per_inch) + y = int(center_y * document.yRes() / points_per_inch) + points.append([x, y]) + + return points + + + def regenerate_lines_from_points(self): + """Regenerate composition lines based on current point positions without calling models""" + document = Krita.instance().activeDocument() + if not document: + print("No active document") + return + + # Read current point positions from the Points layer + current_points = self.read_points_from_layer() + + if not current_points: + print("No points found in Points layer") + return + + if not self.cached_polygon_contours: + print("No cached polygon contours. Please generate adaptive grid first.") + return + + print(f"Found {len(current_points)} points, regenerating lines...") + + # Call the direct function (no server needed) + result = self.regenerate_lines(current_points, self.cached_polygon_contours) + + if not result: + print("Failed to regenerate lines") + return + + self.compose_lines = result['composition_lines'] + print(f"Generated {len(self.compose_lines)} composition lines") + + # Redraw the composition lines + self.draw_composition_lines() + self.update_preview() + self.value_color.append_log_entry("regenerate lines", + f"Regenerated {len(self.compose_lines)} composition lines from {len(current_points)} points") + + def draw_grid(self): + document = Krita.instance().activeDocument() + if document: + w = Krita.instance().activeWindow() + v = w.activeView() + selected_nodes = v.selectedNodes() + + ## get all the rectangles from the selected nodes + custom_rectangles = [] + for node in selected_nodes: + print(f"Selected node: {node.name()} of type {node.type()}") + # Get all rectangles in the vector layer + if node.type() == "vectorlayer": + for shape in node.shapes(): + if shape.type() == "KoPathShape": + # Get the bounding box in points + bbox = shape.boundingBox() + # Convert from points to pixels (72 points per inch) + points_per_inch = 72.0 + x1 = int(bbox.topLeft().x() * document.xRes() / points_per_inch) + y1 = int(bbox.topLeft().y() * document.yRes() / points_per_inch) + x2 = int(bbox.bottomRight().x() * document.xRes() / points_per_inch) + y2 = int(bbox.bottomRight().y() * document.yRes() / points_per_inch) + custom_rectangles.append((x1, y1, x2, y2)) + + ## process the paint layer + reference_layer = document.nodeByName("Reference Image") + if not reference_layer: + print("No reference image found") + return + + width = document.width() + height = document.height() + + # Ensure the temp directory exists and save the reference image + temp_dir = os.getcwd() + if sys.platform == "darwin": # Check if the OS is MacOS + temp_dir = os.path.expanduser("~/Library/Application Support/krita/pykrita/artkrit") + elif sys.platform == "sys": + temp_dir = os.path.expanduser("~/.local/share/krita/pykrita/artkrit") + + temp_dir = os.path.join(temp_dir, "temp") + os.makedirs(temp_dir, exist_ok=True) + temp_path = os.path.join(temp_dir, "krita_temp_image.png") + + print(f"Processing image with file path: {temp_path}") + + # Call the direct processing function (no server needed) + result = self.process_image( + image_path=temp_path, + text_prompt=self.text_prompt_input.text(), + custom_rectangles=custom_rectangles, + polygon_epsilon=self.polygon_epsilon_slider.value() + ) + + if not result: + print("Failed to process image") + return + + print(f"Processing complete: {len(result.get('ploygon_contours', []))} polygons, {len(result.get('points', []))} points, {len(result.get('composition_lines', []))} lines") + + root = document.rootNode() + + ## Draw polygons in krita + ploygon_contours = result['ploygon_contours'] + self.cached_polygon_contours = ploygon_contours # Cache for regeneration + + # Create contours vector layer + contour_layer = document.nodeByName('Contours') + if contour_layer is None: + contour_layer = document.createVectorLayer('Contours') + root.addChildNode(contour_layer, None) + + # Remove existing shapes + for shape in contour_layer.shapes(): + shape.remove() + + document.setActiveNode(contour_layer) + + # Create SVG content for all polygons + svg_content = f'''''' + + for polygon in ploygon_contours: + # Create path data for polygon + points = [f"{point[0]},{point[1]}" for point in polygon] + path_data = "M " + " L ".join(points) + " Z" # Z closes the path + svg_content += f'' + + svg_content += '' + + # Add SVG shapes to the layer + contour_layer.addShapesFromSvg(svg_content) + contour_layer.setVisible(False) + + ## Draw points in krita + points = result['points'] + self.cached_points = points # Cache for reference + + # Create points vector layer + points_layer = document.nodeByName('Points') + if points_layer is None: + points_layer = document.createVectorLayer('Points') + root.addChildNode(points_layer, None) + + # Remove existing shapes + for shape in points_layer.shapes(): + shape.remove() + + document.setActiveNode(points_layer) + + # Create SVG content for all points + svg_content = f'''''' + + for point in points: + x, y = point + # Draw a small circle for each point + svg_content += f'' + + svg_content += '' + + # Add SVG shapes to the layer + points_layer.addShapesFromSvg(svg_content) + points_layer.setVisible(True) + + ## Draw lines in krita + self.compose_lines = result['composition_lines'] + self.draw_composition_lines() + self.update_preview() + self.value_color.append_log_entry("generate adaptive grid", + f'Generating adaptive grid with prompt: {self.text_prompt_input.text()} and polygon epsilon: {self.polygon_epsilon_slider.value()} and number of lines: {self.grid_lines_slider.value()}') + + + def draw_composition_lines(self): + document = Krita.instance().activeDocument() + if not document or len(self.compose_lines) == 0: + return + + width = document.width() + height = document.height() + root = document.rootNode() + + # Create composition vector layer + compose_layer = document.nodeByName('Adaptive Grid') + if compose_layer is None: + compose_layer = document.createVectorLayer('Adaptive Grid') + root.addChildNode(compose_layer, None) + + for shape in compose_layer.shapes(): + shape.remove() + + document.setActiveNode(compose_layer) + + # Create SVG content for all lines + svg_content = f'''''' + num_lines_to_draw = self.grid_lines_slider.value() + num_lines_to_draw = min(num_lines_to_draw, len(self.compose_lines)) + for (i, line) in enumerate(self.compose_lines[:num_lines_to_draw]): + p1, p2 = line + svg_content += f'' + svg_content += '' + + compose_layer.addShapesFromSvg(svg_content) + compose_layer.setVisible(True) + + # Refresh the document + document.refreshProjection() + self.value_color.append_log_entry("draw composition lines", + f"Drew {num_lines_to_draw} composition lines on canvas when asked for {self.grid_lines_slider.value()} lines") + + + def set_reference_image(self): + # First check if there is a reference image layer already + document = Krita.instance().activeDocument() + if not document: + return + + reference_layer = document.nodeByName('Reference Image') + if not reference_layer: + # Open file dialog to select image + file_dialog = QFileDialog() + file_path, _ = file_dialog.getOpenFileName(None, "Select Reference Image", "", "Images (*.png *.jpg *.jpeg *.bmp)") + + if not file_path: + return + + # Get active document + document = Krita.instance().activeDocument() + if not document: + return + + # Read the image + image = QImage(file_path) + if image.isNull(): + return + + # Get document dimensions + doc_width = document.width() + doc_height = document.height() + + # Calculate scaling to fit image within document bounds while preserving aspect ratio + image_aspect = image.width() / image.height() + doc_aspect = doc_width / doc_height + + if image_aspect > doc_aspect: + # Image is wider relative to height - scale to fit width + scaled_width = doc_width + scaled_height = int(doc_width / image_aspect) + else: + # Image is taller relative to width - scale to fit height + scaled_height = doc_height + scaled_width = int(doc_height * image_aspect) + + # Scale the image + scaled_image = image.scaled(scaled_width, scaled_height, Qt.KeepAspectRatio, Qt.SmoothTransformation) + + # Create paint layer + root = document.rootNode() + reference_layer = document.createNode("Reference Image", "paintlayer") + root.addChildNode(reference_layer, None) + + # Calculate position to center the image + x = int((doc_width - scaled_width) / 2) + y = int((doc_height - scaled_height) / 2) + + # Create a temporary QImage with the correct size and format + temp_image = QImage(doc_width, doc_height, QImage.Format_ARGB32) + temp_image.fill(Qt.transparent) + + # Draw the scaled image onto the temporary image + painter = QPainter(temp_image) + painter.drawImage(x, y, scaled_image) + painter.end() + + # Convert to bytes and set as pixel data + ptr = temp_image.bits() + ptr.setsize(temp_image.byteCount()) + byte_array = bytes(ptr) + reference_layer.setPixelData(byte_array, 0, 0, doc_width, doc_height) + + # Refresh document + document.refreshProjection() + + temp_path, half_size_path = self.write_layer_to_temp(reference_layer) + + # Load the image for preview + self.preview_image = QImage(temp_path) + if not self.preview_image.isNull(): + self.update_preview() + self.popout_btn.setEnabled(True) + + self.value_color.upload_image(half_size_path) + self.value_color.append_log_entry("set ref img composition", "Setting reference image for composition") + + + def write_layer_to_temp(self, layer): + document = Krita.instance().activeDocument() + if not document: + return + + # Get active node + node = document.nodeByName(layer.name()) + if not node: + return + + # Read the layer + width = document.width() + height = document.height() + + # Get the layer dimensions + pixel_data = node.pixelData(0, 0, width, height) + + # Create a QImage and copy the pixel data into it + temp_image = QImage(pixel_data, width, height, QImage.Format_RGBA8888).rgbSwapped() + + # Save the reference image to temp directory + temp_dir = os.getcwd() + if sys.platform == "darwin": # Check if the OS is MacOS + temp_dir = os.path.expanduser("~/Library/Application Support/krita/pykrita/artkrit") + elif sys.platform == "sys": + temp_dir = os.path.expanduser("~/.local/share/krita/pykrita/artkrit") + + temp_dir = os.path.join(temp_dir, "temp") + os.makedirs(temp_dir, exist_ok=True) + + # Save the reference image + temp_path = os.path.join(temp_dir, "krita_temp_image.png") + if temp_image.save(temp_path): + print(f"Image saved successfully to {temp_path}") + else: + print("Failed to save image") + + # also write an image half the size + half_size_path = os.path.join(temp_dir, "krita_temp_image_half_size.png") + half_size_image = temp_image.scaled(temp_image.width() // 2, temp_image.height() // 2, Qt.KeepAspectRatio, Qt.SmoothTransformation) + if half_size_image.save(half_size_path): + print(f"Half size image saved successfully to {half_size_path}") + else: + print("Failed to save half size image") + + return temp_path, half_size_path + + + def create_thirds_layer(self): + document = Krita.instance().activeDocument() + if document: + # Create a new transparent layer for the lines + root = document.rootNode() + self.thirds_layer = document.createNode("Rule of Thirds Grid", "paintlayer") + root.addChildNode(self.thirds_layer, None) + + # Get document dimensions + width = document.width() + height = document.height() + + # Create a transparent RGBA image for the layer + image = QImage(width, height, QImage.Format_RGBA8888) + image.fill(Qt.transparent) + + # Draw the lines + painter = QPainter(image) + pen = QPen(QColor(0, 255, 0)) # Green color for lines + pen.setWidth(15) + painter.setPen(pen) + + # Draw vertical lines + for i in range(1, 3): + x = width * i / 3 + painter.drawLine(int(x), 0, int(x), height) + + # Draw horizontal lines + for i in range(1, 3): + y = height * i / 3 + painter.drawLine(0, int(y), width, int(y)) + + painter.end() + + # Convert QImage to bytes and set as pixel data + ptr = image.bits() + ptr.setsize(image.byteCount()) + byte_array = bytes(ptr) + + # Set the pixel data on the layer + self.thirds_layer.setPixelData(byte_array, 0, 0, width, height) + + # Make sure the layer is visible initially + self.thirds_layer.setVisible(True) + + # Refresh the document + document.refreshProjection() + + def create_cross_layer(self): + document = Krita.instance().activeDocument() + if document: + # Create a new transparent layer for the lines + root = document.rootNode() + self.cross_layer = document.createNode("Cross Grid", "paintlayer") + root.addChildNode(self.cross_layer, None) + + # Get document dimensions + width = document.width() + height = document.height() + + # Create a transparent RGBA image for the layer + image = QImage(width, height, QImage.Format_RGBA8888) + image.fill(Qt.transparent) + + # Draw the lines + painter = QPainter(image) + pen = QPen(QColor(0, 255, 0)) # Green color for lines + pen.setWidth(15) + painter.setPen(pen) + + # Draw vertical line (middle) + x = width / 2 + painter.drawLine(int(x), 0, int(x), height) + + # Draw horizontal line (middle) + y = height / 2 + painter.drawLine(0, int(y), width, int(y)) + + painter.end() + + # Convert QImage to bytes and set as pixel data + ptr = image.bits() + ptr.setsize(image.byteCount()) + byte_array = bytes(ptr) + + # Set the pixel data on the layer + self.cross_layer.setPixelData(byte_array, 0, 0, width, height) + + # Make sure the layer is visible initially + self.cross_layer.setVisible(True) + + # Refresh the document + document.refreshProjection() + + def create_circle_layer(self): + document = Krita.instance().activeDocument() + if document: + # Create a new transparent layer for the circle + root = document.rootNode() + self.circle_layer = document.createNode("Circle Grid", "paintlayer") + root.addChildNode(self.circle_layer, None) + + # Get document dimensions + width = document.width() + height = document.height() + + # Create a transparent RGBA image for the layer + image = QImage(width, height, QImage.Format_RGBA8888) + image.fill(Qt.transparent) + + # Draw the circle + painter = QPainter(image) + pen = QPen(QColor(0, 255, 0)) # Green color for circle + pen.setWidth(15) + painter.setPen(pen) + + # Draw a circle in the center + center_x = width / 2 + center_y = height / 2 + radius = min(width, height) / 4 # Adjust radius as needed + painter.drawEllipse(int(center_x - radius), int(center_y - radius), int(radius * 2), int(radius * 2)) + + painter.end() + + # Convert QImage to bytes and set as pixel data + ptr = image.bits() + ptr.setsize(image.byteCount()) + byte_array = bytes(ptr) + + # Set the pixel data on the layer + self.circle_layer.setPixelData(byte_array, 0, 0, width, height) + + # Make sure the layer is visible initially + self.circle_layer.setVisible(True) + + # Refresh the document + document.refreshProjection() + + def toggle_canvas_thirds(self): + document = Krita.instance().activeDocument() + if document: + third_layer = document.nodeByName('Rule of Thirds Grid') + if not third_layer: + self.create_thirds_layer() + self.thirds_visible = True + self.value_color.append_log_entry("toggle rule of thirds grid", "Toggling rule of thirds grid") + else: + # Toggle visibility of the existing layer + current_visibility = third_layer.visible() + third_layer.setVisible(not current_visibility) + self.thirds_visible = not current_visibility + document.refreshProjection() + self.value_color.append_log_entry("toggle rule of thirds grid", "Toggling rule of thirds grid") + self.update_preview() + + def toggle_canvas_cross(self): + document = Krita.instance().activeDocument() + if document: + cross_layer = document.nodeByName('Cross Grid') + if not cross_layer: + self.create_cross_layer() + self.cross_visible = True + self.value_color.append_log_entry("toggle cross grid", "Toggling cross grid") + else: + # Toggle visibility of the existing layer + current_visibility = cross_layer.visible() + cross_layer.setVisible(not current_visibility) + self.cross_visible = not current_visibility + document.refreshProjection() + self.value_color.append_log_entry("toggle cross grid","Toggling cross grid") + self.update_preview() + + def toggle_canvas_circle(self): + document = Krita.instance().activeDocument() + if document: + circle_layer = document.nodeByName('Circle Grid') + if not circle_layer: + self.create_circle_layer() + self.circle_visible = True + self.value_color.append_log_entry("toggle circle grid", "Toggling circle grid") + else: + # Toggle visibility of the existing layer + current_visibility = circle_layer.visible() + circle_layer.setVisible(not current_visibility) + self.circle_visible = not current_visibility + document.refreshProjection() + self.value_color.append_log_entry("toggle circle grid","Toggling circle grid") + self.update_preview() + + def toggle_adaptive_grid(self): + document = Krita.instance().activeDocument() + if document: + adaptive_grid_layer = document.nodeByName('Adaptive Grid') + if adaptive_grid_layer: + # Toggle visibility of the existing layer + current_visibility = adaptive_grid_layer.visible() + adaptive_grid_layer.setVisible(not current_visibility) + self.adaptive_grid_visible = not current_visibility + document.refreshProjection() + self.value_color.append_log_entry("toggle adaptive grid", "Toggling adaptive grid") + self.update_preview() + else: + print("Adaptive grid layer not found") + + def toggle_contours(self): + document = Krita.instance().activeDocument() + if document: + contours_layer = document.nodeByName('Contours') + if contours_layer: + current_visibility = contours_layer.visible() + contours_layer.setVisible(not current_visibility) + self.contours_visible = not current_visibility + document.refreshProjection() + self.value_color.append_log_entry("contours feedback", "Toggling contours visibility") + self.update_preview() + else: + print("Contours layer not found") + + def canvasChanged(self, canvas): + # Reset layer references when canvas changes + self.thirds_layer = None + self.cross_layer = None + self.circle_layer = None + + def krita_sleep(self, value): + loop = QEventLoop() + QTimer.singleShot(value, loop.quit) + loop.exec() + + +# Register the docker with Krita +Krita.instance().addDockWidgetFactory( + DockWidgetFactory("ArtKrit", DockWidgetFactoryBase.DockRight, ArtKrit) +) \ No newline at end of file diff --git a/ArtKrit/requirements.txt b/ArtKrit/requirements.txt new file mode 100644 index 0000000..f79ddad --- /dev/null +++ b/ArtKrit/requirements.txt @@ -0,0 +1,11 @@ +Flask==3.1.0 +flask_cors==5.0.1 +matplotlib==3.10.1 +numpy==2.2.4 +opencv_python==4.11.0.86 +Pillow==11.1.0 +PyQt5==5.15.11 +Requests==2.32.3 +scikit_learn==1.6.1 +transformers==4.49.0 +replicate diff --git a/ArtKrit/script/composition/composition_utils.py b/ArtKrit/script/composition/composition_utils.py new file mode 100644 index 0000000..a23ce2f --- /dev/null +++ b/ArtKrit/script/composition/composition_utils.py @@ -0,0 +1,1010 @@ +import cv2 +import random +import numpy as np +import matplotlib.pyplot as plt +from PIL import Image +from dataclasses import dataclass +from typing import Any, List, Dict, Optional, Union, Tuple +import torch +import requests +import math +from sklearn.neighbors import NearestNeighbors + +## detection result dataclasses +@dataclass +class BoundingBox: + xmin: int + ymin: int + xmax: int + ymax: int + + @property + def xyxy(self) -> List[float]: + return [self.xmin, self.ymin, self.xmax, self.ymax] + +@dataclass +class DetectionResult: + score: float + label: str + box: BoundingBox + mask: Optional[np.array] = None + + @classmethod + def from_dict(cls, detection_dict: Dict) -> 'DetectionResult': + return cls(score=detection_dict['score'], + label=detection_dict['label'], + box=BoundingBox(xmin=detection_dict['box']['xmin'], + ymin=detection_dict['box']['ymin'], + xmax=detection_dict['box']['xmax'], + ymax=detection_dict['box']['ymax'])) + + +def annotate(image: Union[Image.Image, np.ndarray], detection_results: List[DetectionResult], parameters: Dict) -> np.ndarray: + # Convert PIL Image to OpenCV format + image_cv2 = np.array(image) if isinstance(image, Image.Image) else image + image_cv2 = cv2.cvtColor(image_cv2, cv2.COLOR_RGB2BGR) + + ploygon_contours = [] + ploygon_contours_list = [] + + # Iterate over detections and add bounding boxes and masks + for detection in detection_results: + label = detection.label + score = detection.score + box = detection.box + mask = detection.mask + + # Sample a random color for each detection + color = np.random.randint(0, 256, size=3) + + # Draw bounding box + cv2.rectangle(image_cv2, (box.xmin, box.ymin), (box.xmax, box.ymax), color.tolist(), 2) + cv2.putText(image_cv2, f'{label}: {score:.2f}', (box.xmin, box.ymin - 10), cv2.FONT_HERSHEY_SIMPLEX, 1.5, color.tolist(), 3) + + # If mask is available, apply it + if mask is not None: + # Convert mask (various possible dtypes/ranges) to binary uint8 (0/255) + if isinstance(mask, np.ndarray): + m = mask + if m.dtype != np.uint8: + # if in [0,1] float, scale; else cast + if np.max(m) <= 1.0: + m = (m.astype(np.float32) * 255.0).astype(np.uint8) + else: + m = m.astype(np.uint8) + # ensure binary + mask_uint8 = (m > 127).astype(np.uint8) * 255 + else: + # unsupported mask type + print(f"[Annotate] Unsupported mask type: {type(mask)}; skipping") + continue + contours, _ = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE) + approx_contours = [] + approx_contours_list = [] + image_height, image_width = image_cv2.shape[:2] + image_area = float(image_height * image_width) + for contour in contours: + # Filter out tiny blobs and near-full-frame masks + area = cv2.contourArea(contour) + if area < 100: # too small + continue + if area / image_area > 0.85: # too big (likely combined mask) + continue + # Normalize each point in contour by the image width and height + uv_points = np.array([(point[0] / image_width, point[1] / image_height) for point in contour.reshape(-1, 2)], dtype=np.float32) + normalized_arc_length = cv2.arcLength(uv_points, True) + base_eps = parameters['polygon_epsilon'] * normalized_arc_length + eps = max(base_eps, 1e-4) + # Adaptive simplification: avoid 4-corner collapse + approx = cv2.approxPolyDP(uv_points, eps, True) + tries = 0 + while approx.shape[0] <= 6 and eps > 1e-6 and tries < 5: + eps *= 0.5 + approx = cv2.approxPolyDP(uv_points, eps, True) + tries += 1 + # Convert the normalized points back to the original image size + approx_points = np.array([(int(point[0] * image_width), int(point[1] * image_height)) for point in approx.reshape(-1, 2)], dtype=np.int32) + approx_contours.append(approx_points) + approx_contours_list.append(approx_points.reshape(-1, 2).tolist()) + + ploygon_contours.extend(approx_contours) + ploygon_contours_list.extend(approx_contours_list) + if len(approx_contours) > 0: + cv2.drawContours(image_cv2, approx_contours, -1, color.tolist(), 10) + + #### draw composition lines + print("Total polygons: ", len(ploygon_contours)) + + ## get all points + all_points = [] + + ## Find global shortest edge length across all polygons + shortest_edge = float('inf') + polygon_areas = [] + valid_edges = [] # Track all valid edge lengths + + for contour in ploygon_contours: + polygon_areas.append(cv2.contourArea(contour)) + points = contour.reshape(-1, 2) + for i in range(len(points)): + p1 = points[i] + p2 = points[(i + 1) % len(points)] + edge_length = np.linalg.norm(p2 - p1) + if edge_length > 0: # Only consider valid edges + valid_edges.append(edge_length) + shortest_edge = min(shortest_edge, edge_length) + + # FIX: If no valid edges found, use a default based on image size + if not valid_edges or np.isinf(shortest_edge): + image_height, image_width = image_cv2.shape[:2] + shortest_edge = min(image_width, image_height) * 0.01 # 1% of smallest dimension + print(f"[Warning] No valid edges found, using default shortest_edge: {shortest_edge:.2f}") + else: + print(f"[Info] Found shortest_edge: {shortest_edge:.2f} from {len(valid_edges)} valid edges") + + for i, contour in enumerate(ploygon_contours): + # Sample points from the contour edges + sampled_points = sample_contour_points(contour, shortest_edge=shortest_edge) + # Add polygon index to each point + sampled_points_with_index = [(point, i) for point in sampled_points] + all_points.extend(sampled_points_with_index) + + ## merge similar points + print("Total points: ", len(all_points)) + # random.shuffle(all_points) + all_points_with_index = merge_similar_points(all_points, polygon_areas, image_cv2, radius=parameters['point_radius']) + # print("Total points after merging: ", len(all_points)) + + ## find lines that connects at least 4 points + lines = fit_lines(all_points_with_index, image_cv2, line_fit_tol=parameters['line_fit_tol'], inlier_threshold=0.05) + print("Total lines: ", len(lines)) + + ## draw the lines and the points + points_to_draw = [] + for (point, index) in all_points_with_index: + cv2.circle(image_cv2, np.array([point[0], point[1]]).astype(int), 10, (0, 0, 255), -1) + points_to_draw.append([int(point[0]), int(point[1])]) + + lines_list = [] + for line in lines: + ## draw line from the leftmost to the rightmost point + p1, p2 = line_leftmost_to_rightmost(line) + + # img_copy = image_cv2 + # cv2.line(img_copy, np.array(p1).astype(int), np.array(p2).astype(int), (0, 255, 0), 20) + lines_list.append([[int(p1[0]), int(p1[1])], [int(p2[0]), int(p2[1])]]) + + return image_cv2, ploygon_contours_list, lines_list, points_to_draw + +def random_named_css_colors(num_colors: int) -> List[str]: + """ + Returns a list of randomly selected named CSS colors. + + Args: + - num_colors (int): Number of random colors to generate. + + Returns: + - list: List of randomly selected named CSS colors. + """ + # List of named CSS colors + named_css_colors = [ + 'aliceblue', 'antiquewhite', 'aqua', 'aquamarine', 'azure', 'beige', 'bisque', 'black', 'blanchedalmond', + 'blue', 'blueviolet', 'brown', 'burlywood', 'cadetblue', 'chartreuse', 'chocolate', 'coral', 'cornflowerblue', + 'cornsilk', 'crimson', 'cyan', 'darkblue', 'darkcyan', 'darkgoldenrod', 'darkgray', 'darkgreen', 'darkgrey', + 'darkkhaki', 'darkmagenta', 'darkolivegreen', 'darkorange', 'darkorchid', 'darkred', 'darksalmon', 'darkseagreen', + 'darkslateblue', 'darkslategray', 'darkslategrey', 'darkturquoise', 'darkviolet', 'deeppink', 'deepskyblue', + 'dimgray', 'dimgrey', 'dodgerblue', 'firebrick', 'floralwhite', 'forestgreen', 'fuchsia', 'gainsboro', 'ghostwhite', + 'gold', 'goldenrod', 'gray', 'green', 'greenyellow', 'grey', 'honeydew', 'hotpink', 'indianred', 'indigo', 'ivory', + 'khaki', 'lavender', 'lavenderblush', 'lawngreen', 'lemonchiffon', 'lightblue', 'lightcoral', 'lightcyan', 'lightgoldenrodyellow', + 'lightgray', 'lightgreen', 'lightgrey', 'lightpink', 'lightsalmon', 'lightseagreen', 'lightskyblue', 'lightslategray', + 'lightslategrey', 'lightsteelblue', 'lightyellow', 'lime', 'limegreen', 'linen', 'magenta', 'maroon', 'mediumaquamarine', + 'mediumblue', 'mediumorchid', 'mediumpurple', 'mediumseagreen', 'mediumslateblue', 'mediumspringgreen', 'mediumturquoise', + 'mediumvioletred', 'midnightblue', 'mintcream', 'mistyrose', 'moccasin', 'navajowhite', 'navy', 'oldlace', 'olive', + 'olivedrab', 'orange', 'orangered', 'orchid', 'palegoldenrod', 'palegreen', 'paleturquoise', 'palevioletred', 'papayawhip', + 'peachpuff', 'peru', 'pink', 'plum', 'powderblue', 'purple', 'rebeccapurple', 'red', 'rosybrown', 'royalblue', 'saddlebrown', + 'salmon', 'sandybrown', 'seagreen', 'seashell', 'sienna', 'silver', 'skyblue', 'slateblue', 'slategray', 'slategrey', + 'snow', 'springgreen', 'steelblue', 'tan', 'teal', 'thistle', 'tomato', 'turquoise', 'violet', 'wheat', 'white', + 'whitesmoke', 'yellow', 'yellowgreen' + ] + + # Sample random named CSS colors + return random.sample(named_css_colors, min(num_colors, len(named_css_colors))) + +def mask_to_polygon(mask: np.ndarray) -> List[List[int]]: + # Find contours in the binary mask + contours, _ = cv2.findContours(mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + + # Find the contour with the largest area + largest_contour = max(contours, key=cv2.contourArea) + + # Extract the vertices of the contour + polygon = largest_contour.reshape(-1, 2).tolist() + + return polygon + +def polygon_to_mask(polygon: List[Tuple[int, int]], image_shape: Tuple[int, int]) -> np.ndarray: + """ + Convert a polygon to a segmentation mask. + + Args: + - polygon (list): List of (x, y) coordinates representing the vertices of the polygon. + - image_shape (tuple): Shape of the image (height, width) for the mask. + + Returns: + - np.ndarray: Segmentation mask with the polygon filled. + """ + # Create an empty mask + mask = np.zeros(image_shape, dtype=np.uint8) + + # Convert polygon to an array of points + pts = np.array(polygon, dtype=np.int32) + + # Fill the polygon with white color (255) + cv2.fillPoly(mask, [pts], color=(255,)) + + return mask + +def load_image(image_str: str) -> Image.Image: + if image_str.startswith("http"): + image = Image.open(requests.get(image_str, stream=True).raw).convert("RGB") + else: + image = Image.open(image_str).convert("RGB") + return image + +def get_boxes(results: DetectionResult) -> List[List[List[float]]]: + boxes = [] + for result in results: + xyxy = result.box.xyxy + boxes.append(xyxy) + + return [boxes] + +def refine_masks(masks: torch.BoolTensor, polygon_refinement: bool = False) -> List[np.ndarray]: + masks = masks.cpu().float() + masks = masks.permute(0, 2, 3, 1) + masks = masks.mean(axis=-1) + masks = (masks > 0).int() + masks = masks.numpy().astype(np.uint8) + masks = list(masks) + + if polygon_refinement: + for idx, mask in enumerate(masks): + shape = mask.shape + polygon = mask_to_polygon(mask) + mask = polygon_to_mask(polygon, shape) + masks[idx] = mask + + return masks + +def group_consecutive(numbers): + res = [] + for i in range(len(numbers) - 1): + res.append([numbers[i], numbers[i + 1]]) + return res + +## first attempt to draw composition lines +def lines_from_collinear_edges(ploygon_contours, image_cv2): + ## finding collinear lines + ## not comparing against itself + ## lines from different detections might still overlap with each other since the detections might overlap + for i in range(len(ploygon_contours)): + for j in range(i+1, len(ploygon_contours)): + # print(len(ploygon_contours[i]), len(ploygon_contours[j])) + lines_i = group_consecutive(ploygon_contours[i].reshape(-1, 2).tolist()) + lines_j = group_consecutive(ploygon_contours[j].reshape(-1, 2).tolist()) + counter = 0 + for line_i in lines_i: + for line_j in lines_j: + if are_lines_collinear(line_i, line_j, parallel_tol=1e-4, col_tol=0.05) and not are_lines_copoint(line_i, line_j, tol=10): + counter += 1 + # plot the lines + cv2.line(image_cv2, (line_i[0][0], line_i[0][1]), (line_i[1][0], line_i[1][1]), (0, 255, 0), 40) + cv2.line(image_cv2, (line_j[0][0], line_j[0][1]), (line_j[1][0], line_j[1][1]), (0, 0, 255), 30) + print("Collinear") + print(line_i, line_j) + pt_1, pt_2 = average_lines(line_i, line_j) + + pts = extend_line_to_edge(image_cv2.shape[:2], (pt_1, pt_2)) + + cv2.line(image_cv2, (pts[0][0], pts[0][1]), (pts[1][0], pts[1][1]), (255, 0, 0), 30) + + +def are_lines_copoint(line1, line2, tol=1e-9): + """ + Check if two lines are copoint + + Args: + tol: in pixels + + Returns: + - True if the lines are copoint, False otherwise. + """ + p1, p2 = np.array(line1[0]), np.array(line1[1]) + q1, q2 = np.array(line2[0]), np.array(line2[1]) + + dist_p1_q1 = np.linalg.norm(p1 - q1) + dist_p1_q2 = np.linalg.norm(p1 - q2) + dist_p2_q1 = np.linalg.norm(p2 - q1) + dist_p2_q2 = np.linalg.norm(p2 - q2) + + if dist_p1_q1 < tol or dist_p2_q2 < tol or dist_p1_q2 < tol or dist_p2_q1 < tol: + return True + + return False + + +def are_lines_collinear(line1, line2, parallel_tol=1e-9, col_tol=1e-9): + """ + Check if two lines are collinear in 2D or 3D space. + + Parameters: + - line1: Tuple of two points defining the first line (e.g., ((x1, y1), (x2, y2)) or ((x1, y1, z1), (x2, y2, z2))). + - line2: Tuple of two points defining the second line. + - tolerance: A small value to account for floating-point inaccuracies. + + Returns: + - True if the lines are collinear, False otherwise. + """ + # Extract points from the input + p1, p2 = np.array(line1[0]), np.array(line1[1]) + q1, q2 = np.array(line2[0]), np.array(line2[1]) + + # Compute direction vectors of both lines + dir1 = (p2 - p1) / np.linalg.norm(p2 - p1) + dir2 = (q2 - q1) / np.linalg.norm(q2 - q1) + + # Check if direction vectors are parallel using the cross product + dot_product = np.dot(dir1, dir2) + if not np.isclose(dot_product, 1.0, atol=parallel_tol): + return False # Lines are not parallel + + # Check if a point from one line lies on the other line + # Vector from a point on line1 to a point on line2 + # get the closer point to p1 + vector_between_lines = q2 - p1 + if np.linalg.norm(q1 - p1) < np.linalg.norm(q2 - p1): + vector_between_lines = q1 - p1 + + dir_between = vector_between_lines / np.linalg.norm(vector_between_lines) + + # dot product of this vector with one of the direction vectors + alignment_check = np.dot(dir1, dir_between) + if not np.isclose(alignment_check, 1.0, atol=col_tol): + return False # A point from one line does not lie on the other + + return True # Lines are collinear + + +def average_lines(line1, line2): + """ + Average a set of lines in 2D or 3D space. + + Parameters: + - lines: List of tuples, each containing two points defining a line. + + Returns: + - Tuple of two points representing the average line. + """ + # Extract points from the input + p1, p2 = np.array(line1[0]), np.array(line1[1]) + q1, q2 = np.array(line2[0]), np.array(line2[1]) + + pt_1 = (p1 + q1) / 2 + pt_2 = (p2 + q2) / 2 + + return pt_1.astype(int), pt_2.astype(int) + + +def extend_line_to_edge(dimensions, line, SCALE=10): + """ + Based on https://stackoverflow.com/questions/72083896/how-to-stretch-a-line-to-fit-image-with-python-opencv + """ + p1 = line[0] + p2 = line[1] + + # Calculate the intersection point given (x1, y1) and (x2, y2) + def line_intersection(line1, line2): + x_diff = (line1[0][0] - line1[1][0], line2[0][0] - line2[1][0]) + y_diff = (line1[0][1] - line1[1][1], line2[0][1] - line2[1][1]) + + def detect(a, b): + return a[0] * b[1] - a[1] * b[0] + + div = detect(x_diff, y_diff) + if div == 0: + raise Exception('lines do not intersect') + + dist = (detect(*line1), detect(*line2)) + x = detect(dist, x_diff) / div + y = detect(dist, y_diff) / div + return int(x), int(y) + + x1, x2 = 0, 0 + y1, y2 = 0, 0 + + # Extract w and h regardless of grayscale or BGR image + if len(dimensions) == 3: + h, w, _ = dimensions + elif len(dimensions) == 2: + h, w = dimensions + + # Take longest dimension and use it as maxed out distance + if w > h: + distance = SCALE * w + else: + distance = SCALE * h + + # Reorder smaller X or Y to be the first point + # and larger X or Y to be the second point + try: + slope = (p2[1] - p1[1]) / (p1[0] - p2[0]) + # HORIZONTAL or DIAGONAL + if p1[0] <= p2[0]: + x1, y1 = p1 + x2, y2 = p2 + else: + x1, y1 = p2 + x2, y2 = p1 + except ZeroDivisionError: + # VERTICAL + if p1[1] <= p2[1]: + x1, y1 = p1 + x2, y2 = p2 + else: + x1, y1 = p2 + x2, y2 = p1 + + # Extend after end-point A + length_A = math.sqrt((x2 - x1)**2 + (y2 - y1)**2) + p3_x = int(x1 + (x1 - x2) / length_A * distance) + p3_y = int(y1 + (y1 - y2) / length_A * distance) + + # Extend after end-point B + length_B = math.sqrt((x1 - x2)**2 + (y1 - y2)**2) + p4_x = int(x2 + (x2 - x1) / length_B * distance) + p4_y = int(y2 + (y2 - y1) / length_B * distance) + + # -------------------------------------- + # Limit coordinates to borders of image + # -------------------------------------- + # HORIZONTAL + if y1 == y2: + if p3_x < 0: + p3_x = 0 + if p4_x > w: + p4_x = w + return ((p3_x, p3_y), (p4_x, p4_y)) + # VERTICAL + elif x1 == x2: + if p3_y < 0: + p3_y = 0 + if p4_y > h: + p4_y = h + return ((p3_x, p3_y), (p4_x, p4_y)) + # DIAGONAL + else: + A = (p3_x, p3_y) + B = (p4_x, p4_y) + + C = (0, 0) # C-------D + D = (w, 0) # |-------| + E = (w, h) # |-------| + F = (0, h) # F-------E + + if slope > 0: + # 1st point, try C-F side first, if OTB then F-E + new_x1, new_y1 = line_intersection((A, B), (C, F)) + if new_x1 > w or new_y1 > h: + new_x1, new_y1 = line_intersection((A, B), (F, E)) + + # 2nd point, try C-D side first, if OTB then D-E + new_x2, new_y2 = line_intersection((A, B), (C, D)) + if new_x2 > w or new_y2 > h: + new_x2, new_y2 = line_intersection((A, B), (D, E)) + + return ((new_x1, new_y1), (new_x2, new_y2)) + elif slope < 0: + # 1st point, try C-F side first, if OTB then C-D + new_x1, new_y1 = line_intersection((A, B), (C, F)) + if new_x1 < 0 or new_y1 < 0: + new_x1, new_y1 = line_intersection((A, B), (C, D)) + # 2nd point, try F-E side first, if OTB then E-D + new_x2, new_y2 = line_intersection((A, B), (F, E)) + if new_x2 > w or new_y2 > h: + new_x2, new_y2 = line_intersection((A, B), (E, D)) + return ((new_x1, new_y1), (new_x2, new_y2)) + + +# get the y=mx+c equation from given points +def get_slope_and_intercept(pointA, pointB): + slope = (pointB[1] - pointA[1])/(pointB[0] - pointA[0]) + intercept = pointB[1] - slope * pointB[0] + return slope, intercept + +def sample_contour_points(contour, shortest_edge=1): + """ + Sample points from a contour's edges with density proportional to edge length. + + Args: + contour: OpenCV contour (numpy array of points) + shortest_edge: Reference edge length for sampling density + + Returns: + List of sampled points [[x1,y1], [x2,y2], ...] + """ + # Safety check: validate inputs + if len(contour) < 2: + print(f"[Warning] Contour has only {len(contour)} points, returning as-is") + return contour.reshape(-1, 2).tolist() + + if shortest_edge <= 0 or np.isinf(shortest_edge) or np.isnan(shortest_edge): + print(f"[Warning] Invalid shortest_edge value: {shortest_edge}, using default") + shortest_edge = 10.0 # Fallback default + + points = contour.reshape(-1, 2) + sampled_points = [] + + # Iterate through each edge of the contour + for i in range(len(points)): + # Get current and next point (handle wrap-around) + p1 = points[i] + p2 = points[(i + 1) % len(points)] + + # Calculate edge length + edge_length = np.linalg.norm(p2 - p1) + + # Skip zero-length edges + if edge_length < 1e-6: + continue + + # Calculate number of points to sample for this edge + num_points = max(2, int(edge_length // (shortest_edge * 2))) + + # Additional safety check + if num_points > 10000: # Prevent excessive sampling + num_points = 10000 + + # Sample points along the edge + t = np.linspace(0, 1, num_points) + for t_val in t: + # Linear interpolation between p1 and p2 + x = p1[0] + t_val * (p2[0] - p1[0]) + y = p1[1] + t_val * (p2[1] - p1[1]) + sampled_points.append([x, y]) + + # If no points were sampled, return the original contour points + if len(sampled_points) == 0: + print("[Warning] No points sampled, returning original contour") + return points.tolist() + + return sampled_points + +# a simple algorithm to find the hash of slope and intercept +def get_unique_id(slope, intercept): + return str(slope)+str(intercept) + + +def exists_slope_intercept(slope, intercept, slope_intercepts, s_tol=1e-4, i_tol=1e-4): + for slope_intercept in slope_intercepts: + if math.isclose(slope, slope_intercept[0], abs_tol=s_tol) and math.isclose(intercept, slope_intercept[1], abs_tol=i_tol): + return True + return False + + +def fit_lines(all_points_with_index, image_cv2, line_fit_tol=1, inlier_threshold=0.1): + """ + Fit lines to points using RANSAC algorithm. + + Args: + all_points: List of points [[x1,y1], [x2,y2], ...] + line_fit_tol: Tolerance for point-to-line distance + min_points_per_line: Minimum number of points required to form a line + ransac_iterations: Number of RANSAC iterations + inlier_threshold: Fraction of points that need to be inliers to consider a line valid + + Returns: + List of lines, where each line is a list of points that fit that line + """ + min_points_per_line = 4 + if len(all_points_with_index) < min_points_per_line: + return [] + + lines = [] + all_points = [pt[0] for pt in all_points_with_index] + polygon_indices = [pt[1] for pt in all_points_with_index] + + ## convert all points to uv coordinates + remaining_points = np.array([[pt[0] / image_cv2.shape[1], pt[1] / image_cv2.shape[0]] for pt in all_points]) + remaining_indices = np.array(polygon_indices) + + p_success = 0.99 + w = inlier_threshold # probability of selecting an inlier + sample_size = 2 + k = int(np.ceil(np.log(1 - p_success) / np.log(1 - w**sample_size))) + ransac_iterations = 2*k # Use minimum between computed and provided iterations + + while len(remaining_points) >= min_points_per_line: + best_line = None + best_inliers = None + best_inlier_count = 0 + + # RANSAC iterations + # Compute required number of RANSAC iterations based on: + # - probability of success (0.99) + # - inlier ratio (inlier_threshold) + # - number of points needed for model (2) + + for _ in range(ransac_iterations): + # Get unique polygon indices + unique_polygons = np.unique(remaining_indices) + if len(unique_polygons) < 2: + continue + + # Pick two different random polygons + poly1, poly2 = np.random.choice(unique_polygons, 2, replace=False) + + # Get points from first polygon + points1 = remaining_points[remaining_indices == poly1] + if len(points1) == 0: + continue + p1 = points1[np.random.randint(len(points1))] + + # Get points from second polygon + points2 = remaining_points[remaining_indices == poly2] + if len(points2) == 0: + continue + p2 = points2[np.random.randint(len(points2))] + + # Skip if points are too close + if np.allclose(p1, p2): + continue + + # Get line parameters (ax + by + c = 0) + line_vector = p2 - p1 + line_vector = line_vector / np.linalg.norm(line_vector) + a, b = -line_vector[1], line_vector[0] # normal vector + c = -(a * p1[0] + b * p1[1]) + + # Calculate distances from all points to the line + distances = np.abs(a * remaining_points[:, 0] + b * remaining_points[:, 1] + c) + inliers = distances < line_fit_tol + + inlier_count = np.sum(inliers) + + if inlier_count > best_inlier_count: + best_line = (a, b, c) + best_inliers = inliers + best_inlier_count = inlier_count + + + # Check if we found a good line + if best_inlier_count >= min_points_per_line and best_inlier_count / len(all_points) >= inlier_threshold: + print(f"Found a good line with {best_inlier_count} inliers ({best_inlier_count/len(all_points)*100:.1f}%)") + + # Add the line and its inliers to our results + line_points = remaining_points[best_inliers].tolist() + line_points = [[pt[0] * image_cv2.shape[1], pt[1] * image_cv2.shape[0]] for pt in line_points] + lines.append(line_points) + + # Remove the inliers from remaining points + remaining_points = remaining_points[~best_inliers] + remaining_indices = remaining_indices[~best_inliers] + else: + # No good line found, stop + break + + return lines + + +def merge_similar_points(points_with_index, polygon_areas, image, radius=0.00001): + if len(points_with_index) == 0: + return np.array([]) + + image_width = image.shape[1] + image_height = image.shape[0] + + points = [pt[0] for pt in points_with_index] + polygon_indices = [pt[1] for pt in points_with_index] + + ## normalize the points + point_features = [[pt[0] / image_width, pt[1] / image_height] for pt in points] + + neighbors = NearestNeighbors(radius=radius) + neighbors.fit(point_features) + + distances, indices = neighbors.radius_neighbors(point_features) + print("Total point indices: ", len(indices)) + + print("Total point duplicates: ", len([i for i in indices if len(i) > 1])) + + dup_indices = set() + unique_points = [] + for i in range(len(points_with_index)): + if i in dup_indices: + continue + + if len(indices[i]) > 1: + cluster_points = [points[c_i] for c_i in indices[i]] + cluster_polygon_areas = [polygon_areas[polygon_indices[c_i]] for c_i in indices[i]] + max_area_idx = np.argmax(cluster_polygon_areas) + polygon_index = polygon_indices[indices[i][max_area_idx]] + cluster_points = np.array(cluster_points) + cluster_center = np.mean(cluster_points, axis=0) + unique_points.append((cluster_center, polygon_index)) + for c_i in indices[i]: + dup_indices.add(c_i) + else: + unique_points.append((points[i], polygon_indices[i])) + + return unique_points + + +def merge_similar_lines(lines, image, radius=1): + if len(lines) == 0: + return np.array([]) + + image_width = image.shape[1] + image_height = image.shape[0] + + line_features = [] + for i, line in enumerate(lines): + ## Fit a line through the points in 'line' + # [vx, vy, x, y] = cv2.fitLine(np.array(line), cv2.DIST_L2, 0, 0.01, 0.01) + # lefty = int((-x * vy / vx) + y) + # righty = int(((image.shape[1] - x) * vy / vx) + y) + p1, p2 = line_leftmost_to_rightmost(line) + + ## normalize the points + line_features.append([p1[0] / image_width, p1[1] / image_height, p2[0] / image_width, p2[1] / image_height]) + + neighbors = NearestNeighbors(algorithm='ball_tree', radius=radius, metric=line_segment_metric) + neighbors.fit(line_features) + distances, indices = neighbors.radius_neighbors(line_features) + print("Total line indices: ", len(indices)) + print("Total line duplicates: ", len([i_line for i_line in indices if len(i_line) > 1])) + + dup_indices = set() + unique_lines = [] + for i in range(len(lines)): + if i in dup_indices: + continue + + if len(indices[i]) > 1: + line_points = [] + for c_i in indices[i]: + line_points.extend(lines[c_i]) + dup_indices.add(c_i) + unique_lines.append(line_points) + else: + unique_lines.append(lines[i]) + + print("Total unique lines: ", len(unique_lines)) + return unique_lines + + +def line_segment_metric(line1, line2): + p1, p2 = np.array([line1[0], line1[1]]), np.array((line1[2], line1[3])) + q1, q2 = np.array([line2[0], line2[1]]), np.array((line2[2], line2[3])) + + dir1 = (p2 - p1) / np.linalg.norm(p2 - p1) + dir2 = (q2 - q1) / np.linalg.norm(q2 - q1) + + ## parallel + dot_product = np.dot(dir1, dir2) + + ## line segment distance (line points should have already been normalized) + distance = segments_distance(line1[0], line1[1], line1[2], line1[3], line2[0], line2[1], line2[2], line2[3]) + + return (1 - dot_product) + distance + + +## three functions below are taken from +## https://stackoverflow.com/questions/2824478/shortest-distance-between-two-line-segments +def segments_distance(x11, y11, x12, y12, x21, y21, x22, y22): + """ distance between two segments in the plane: + one segment is (x11, y11) to (x12, y12) + the other is (x21, y21) to (x22, y22) + """ + if segments_intersect(x11, y11, x12, y12, x21, y21, x22, y22): + return 0 + + # try each of the 4 vertices w/the other segment + distances = [] + distances.append(point_segment_distance(x11, y11, x21, y21, x22, y22)) + distances.append(point_segment_distance(x12, y12, x21, y21, x22, y22)) + distances.append(point_segment_distance(x21, y21, x11, y11, x12, y12)) + distances.append(point_segment_distance(x22, y22, x11, y11, x12, y12)) + return min(distances) + + +def segments_intersect(x11, y11, x12, y12, x21, y21, x22, y22): + """ whether two segments in the plane intersect: + one segment is (x11, y11) to (x12, y12) + the other is (x21, y21) to (x22, y22) + """ + dx1 = x12 - x11 + dy1 = y12 - y11 + dx2 = x22 - x21 + dy2 = y22 - y21 + delta = dx2 * dy1 - dy2 * dx1 + if delta == 0: return False # parallel segments + s = (dx1 * (y21 - y11) + dy1 * (x11 - x21)) / delta + t = (dx2 * (y11 - y21) + dy2 * (x21 - x11)) / (-delta) + return (0 <= s <= 1) and (0 <= t <= 1) + + +def point_segment_distance(px, py, x1, y1, x2, y2): + dx = x2 - x1 + dy = y2 - y1 + if dx == dy == 0: # the segment's just a point + return math.hypot(px - x1, py - y1) + + # Calculate the t that minimizes the distance. + t = ((px - x1) * dx + (py - y1) * dy) / (dx * dx + dy * dy) + + # See if this represents one of the segment's + # end points or a point in the middle. + if t < 0: + dx = px - x1 + dy = py - y1 + elif t > 1: + dx = px - x2 + dy = py - y2 + else: + near_x = x1 + t * dx + near_y = y1 + t * dy + dx = px - near_x + dy = py - near_y + + return math.hypot(dx, dy) + + +def line_leftmost_to_rightmost(line): + """_summary_ + + Args: + line (list): a list of points in the form of [[x1, y1], [x2, y2], ...] + """ + [vx, vy, x, y] = cv2.fitLine(np.array(line), cv2.DIST_L2, 0, 0.01, 0.01) + points = sorted(line, key=lambda point: (point[0], point[1])) + leftmost_x = points[0][0] + rightmost_x = points[-1][0] + slope = float(vy / vx + 0.00001) + intercept = float(points[0][1]) - slope * float(points[0][0]) + leftmost_y = slope * leftmost_x + intercept + rightmost_y = slope * rightmost_x + intercept + return (leftmost_x, leftmost_y), (rightmost_x, rightmost_y) + +def process_image_direct(image, detections, polygon_epsilon): + """ + Direct processing without server - ANNOTATION ONLY. + Detection and segmentation should be done by the caller. + + Args: + image: PIL Image object + detections: List of DetectionResult objects (already detected and segmented) + polygon_epsilon: Epsilon for polygon approximation + + Returns: + Dictionary with polygon_contours, composition_lines, and points + """ + import cv2 + import time + + t0 = time.time() + + # Annotate + image_array = np.array(image) + + visualization_parameters = { + "polygon_epsilon": polygon_epsilon * 1e-3, + "point_radius": 1e-2, + "line_fit_tol": 0.04, + "line_radius": 1e-1 + } + + annotated_image, polygon_contours_list, lines_list, points_to_draw = annotate( + image_array, detections, visualization_parameters + ) + t_annotate = time.time() + + # Save annotated image + cv2.imwrite("temp/krita_temp_detection_res.png", img=annotated_image) + + # Timing summary + try: + print(f"[Timing] annotate={t_annotate - t0:.2f}s total={t_annotate - t0:.2f}s") + except Exception as e: + print(f"[Timing] error computing timings: {e}") + + result = { + "ploygon_contours": polygon_contours_list, + "composition_lines": lines_list, + "points": points_to_draw + } + + return result + +def regenerate_lines_direct(points, polygon_contours): + """ + Regenerate composition lines from manually adjusted points + without calling the detection models again. + + Args: + points: List of [x, y] coordinates + polygon_contours: List of polygon contours (each is a list of [x, y] coordinates) + + Returns: + List of composition lines [[[x1, y1], [x2, y2]], ...] + """ + import time + + if not points: + raise ValueError('Points are required') + + if not polygon_contours: + raise ValueError('Polygon contours are required') + + print(f"[Direct] Regenerating lines from {len(points)} manually adjusted points") + t0 = time.time() + + # Assign points to polygons + points_with_index = assign_points_to_polygons(points, polygon_contours) + + # Create a dummy image array for shape information + max_x = max(max(p[0] for p in contour) for contour in polygon_contours) + max_y = max(max(p[1] for p in contour) for contour in polygon_contours) + image_shape = (int(max_y) + 1, int(max_x) + 1, 3) + dummy_image = np.zeros(image_shape, dtype=np.uint8) + + # Use the line fitting logic + line_fit_tol = 0.04 + inlier_threshold = 0.05 + lines = fit_lines(points_with_index, dummy_image, line_fit_tol=line_fit_tol, inlier_threshold=inlier_threshold) + + t_generate = time.time() + print(f"[Direct] Generated {len(lines)} composition lines in {t_generate - t0:.2f}s") + + # Convert lines to the same format as process_image + lines_list = [] + for line in lines: + p1, p2 = line_leftmost_to_rightmost(line) + lines_list.append([[int(p1[0]), int(p1[1])], [int(p2[0]), int(p2[1])]]) + + return lines_list + + +def assign_points_to_polygons(points, polygon_contours): + """ + Assign each point to the polygon it's closest to. + + Args: + points: List of [x, y] coordinates + polygon_contours: List of polygon contours (each is a list of [x, y] coordinates) + + Returns: + List of tuples: [(point, polygon_index), ...] + """ + points_with_index = [] + + for point in points: + point_array = np.array(point, dtype=np.float32) + min_distance = float('inf') + closest_polygon_idx = 0 + + # Find the closest polygon to this point + for poly_idx, contour in enumerate(polygon_contours): + contour_array = np.array(contour, dtype=np.float32).reshape(-1, 2) + + # Calculate distance to each point in the contour + distances = np.linalg.norm(contour_array - point_array, axis=1) + min_dist_to_contour = np.min(distances) + + if min_dist_to_contour < min_distance: + min_distance = min_dist_to_contour + closest_polygon_idx = poly_idx + + points_with_index.append((point, closest_polygon_idx)) + + print(f"[Direct] Assigned {len(points)} points to {len(polygon_contours)} polygons") + return points_with_index \ No newline at end of file diff --git a/ArtKrit/script/composition/run_models.py b/ArtKrit/script/composition/run_models.py new file mode 100644 index 0000000..eef3e95 --- /dev/null +++ b/ArtKrit/script/composition/run_models.py @@ -0,0 +1,545 @@ +from typing import Any, Dict, List +from PIL import Image +from io import BytesIO +import requests +import numpy as np +from PIL import Image + + +import torch +import replicate + +from .composition_utils import * + +try: + from replicate.helpers import FileOutput as ReplicateFileOutput +except Exception: + ReplicateFileOutput = None + +# ------------------------------ +# Model versions +# ------------------------------ +REPLICATE_MODEL = "adirik/grounding-dino:efd10a8ddc57ea28773327e881ce95e20cc1d734c589f7dd01d2036921ed78aa" +REPLICATE_SAM_MODEL = "meta/sam-2:fe97b453a6455861e3bac769b441ca1f1086110da7466dbb65cf1eecfd60dc83" + +# ------------------------------ +# Detector +# ------------------------------ +class ReplicateGroundingDetector: + def __init__(self, model_version: str = REPLICATE_MODEL, box_threshold: float = 0.2, text_threshold: float = 0.2): + self.model_version = model_version + self.box_threshold = box_threshold + self.text_threshold = text_threshold + + def __call__(self, image: Image.Image, candidate_labels: List[str], threshold: float = None) -> List[Dict[str, Any]]: + import base64 + + # Prepare query + labels = [l if l.endswith(".") else (l + ".") for l in candidate_labels] + query = " ".join(labels) + + # OPTIMIZATION: Downscale image before uploading to reduce payload size + max_side = 1280 # Replicate's GroundingDINO works well at this resolution + W, H = image.size + if max(W, H) > max_side: + scale = max_side / max(W, H) + new_size = (int(W * scale), int(H * scale)) + image_upload = image.resize(new_size, Image.BILINEAR) + print(f"[GroundingDINO] Downscaling for upload: {(W, H)} -> {new_size}") + else: + image_upload = image + scale = 1.0 + + # Convert image to base64 data URI with JPEG (smaller than PNG) + buffered = BytesIO() + image_upload.save(buffered, format="JPEG", quality=85) + img_str = base64.b64encode(buffered.getvalue()).decode() + data_uri = f"data:image/jpeg;base64,{img_str}" + + # Replicate input + inputs = { + "image": data_uri, + "query": query, + "box_threshold": threshold if threshold is not None else self.box_threshold, + "text_threshold": threshold if threshold is not None else self.text_threshold, + } + + print(f"[Replicate][GroundingDINO] model={self.model_version}\n query=\"{query}\"\n box_threshold={inputs['box_threshold']} text_threshold={inputs['text_threshold']}") + print(f"[Replicate][GroundingDINO] payload size: {len(img_str) / 1024:.1f} KB") + + # Run model with extended timeout + try: + out = replicate.run(self.model_version, input=inputs) + except Exception as e: + print(f"[Replicate][GroundingDINO] Error: {e}") + raise + + # Parse detections and scale boxes back to original size + detections = out.get("detections", []) + print(f"[Replicate][GroundingDINO] raw detections: {len(detections)}") + results = [] + for d in detections: + bbox = d["bbox"] + item = { + "score": d.get("score", 0.0), + "label": d.get("label", ""), + "box": { + "xmin": int(bbox[0] / scale), + "ymin": int(bbox[1] / scale), + "xmax": int(bbox[2] / scale), + "ymax": int(bbox[3] / scale), + }, + } + results.append(item) + if results: + print("[Replicate][GroundingDINO] parsed detections:") + for r in results: + b = r["box"] + print(f" - {r['label']} score={r['score']:.2f} box=[{b['xmin']:.1f},{b['ymin']:.1f},{b['xmax']:.1f},{b['ymax']:.1f}]") + return results + +def download_mask(url): + resp = requests.get(url) + img = Image.open(BytesIO(resp.content)) + # Prefer alpha channel if present (many mask PNGs encode mask in alpha) + if "A" in img.getbands(): + mask_pil = img.getchannel("A") + else: + mask_pil = img.convert("L") + mask_np = np.array(mask_pil) + return mask_np + + +# ------------------------------ +# Cloud SAM wrapper +# ------------------------------ +def replicate_sam(image_file, boxes=None, **kwargs): + """ + Cloud-based SAM segmentation using Replicate. + Accepts a local file-like object (BytesIO or open file). + If `boxes` is provided, attempt box-prompted segmentation (xyxy in pixel coords). + """ + # OPTIMIZED: Much lighter configuration to prevent timeouts + inputs = { + "image": image_file, + "use_m2m": False, + # CRITICAL: Reduce points_per_side dramatically (default is 32!) + "points_per_side": 4, # Very aggressive reduction + # Stricter thresholds = fewer masks + "pred_iou_thresh": 0.90, + "stability_score_thresh": 0.92, + # Additional optimizations + "crop_n_layers": 0, # Disable crop-based refinement + "crop_n_points_downscale_factor": 2, + } + # Try common box prompt field names used by Replicate SAM variants + if boxes is not None: + inputs["bboxes"] = boxes + inputs["input_boxes"] = boxes + inputs["boxes"] = boxes + inputs["box_format"] = "xyxy" + inputs["return_individual_masks"] = True + inputs.update(kwargs) + print(f"[Replicate][SAM] model={REPLICATE_SAM_MODEL} calling with keys={list(inputs.keys())}") + + try: + out = replicate.run(REPLICATE_SAM_MODEL, input=inputs) + try: + if isinstance(out, dict): + print(f"[Replicate][SAM] returned type={type(out)} keys={list(out.keys())}") + elif isinstance(out, (list, tuple)): + print(f"[Replicate][SAM] returned {len(out)} items") + else: + print(f"[Replicate][SAM] returned type={type(out)}") + except Exception as e: + print(f"[Replicate][SAM] logging error: {e}") + return out + except Exception as e: + print(f"[Replicate][SAM] Error: {e}") + raise + +# ------------------------------ +# Initialize models +# ------------------------------ +def init_models(): + """ + Initialize the object detector (Replicate Grounding DINO) + and cloud-based SAM segmenter (Replicate SAM). + """ + # No local device needed; everything is on Replicate + object_detector = ReplicateGroundingDetector( + model_version=REPLICATE_MODEL, + box_threshold=0.2, + text_threshold=0.2, + ) + + # Cloud SAM: just a function reference + segmentator = replicate_sam + processor = None # not needed for cloud SAM + + print("✅ Initialized: Using Replicate for detection and segmentation") + print(f" - Detector: {REPLICATE_MODEL}") + print(f" - SAM: {REPLICATE_SAM_MODEL}") + return object_detector, segmentator, processor + +# ------------------------------ +# Detection pipeline +# ------------------------------ +def detect( + image: Image.Image, + labels: List[str], + detector: ReplicateGroundingDetector, + threshold: float = 0.3, +) -> List[Dict[str, Any]]: + """ + Detect objects with Replicate Grounding DINO and filter overly large boxes. + """ + print(f"[Detect] labels={labels} threshold={threshold}") + raw_results = detector(image, candidate_labels=labels, threshold=threshold) + image_area = image.size[0] * image.size[1] + + filtered_results = [] + for r in raw_results: + xmin, ymin, xmax, ymax = r["box"]["xmin"], r["box"]["ymin"], r["box"]["xmax"], r["box"]["ymax"] + box_area = (xmax - xmin) * (ymax - ymin) + if image_area > 0 and (box_area / image_area) < 0.8: + filtered_results.append(DetectionResult.from_dict(r)) + print(f"[Detect] raw={len(raw_results)} filtered={len(filtered_results)}") + for d in filtered_results: + print(f"[Detect] keep {d.label} score={d.score:.2f} box={d.box.xyxy}") + return filtered_results + +# ------------------------------ +# Segmentation pipeline +# ------------------------------ + +def _parse_sam_output(out): + """Normalize various possible Replicate SAM outputs to a list.""" + if isinstance(out, dict): + print(f"[Segment] SAM dict keys: {list(out.keys())}") + # Special-case common schema from meta/sam-2 on Replicate + if "combined_mask" in out: + ims = out.get("individual_masks") + parsed = [] + # prefer individual masks first + if isinstance(ims, list): + for m in ims: + if isinstance(m, dict) and "mask" in m: + parsed.append(m["mask"]) # unwrap inner mask field + else: + parsed.append(m) + # then include combined if present + cm = out.get("combined_mask") + if isinstance(cm, dict) and "mask" in cm: + parsed.append(cm["mask"]) # unwrap + elif cm is not None: + parsed.append(cm) + return parsed + for key in ["masks", "mask", "segments", "segmentations", "output", "data"]: + if key in out: + return out[key] if isinstance(out[key], list) else [out[key]] + # fallback: first list-like value + for v in out.values(): + if isinstance(v, list): + return v + return [] + if isinstance(out, (list, tuple)): + return list(out) + return [out] + + +def _coerce_mask_to_numpy(mask_data, target_hw): + """ + Convert mask outputs (url, data-uri, PIL, ndarray, dict) to a HxW uint8 binary numpy array (0 or 255). + target_hw = (H, W) of the original image; masks will be resized to this. + """ + import base64 + H, W = target_hw + + def _ensure_size_u8(m): + # squeeze channel if needed + if m.ndim == 3: + m = m[..., 0] + if m.shape != (H, W): + pil = Image.fromarray(m) + pil = pil.resize((W, H), resample=Image.NEAREST) + m = np.array(pil) + # normalize to uint8 binary + if m.dtype != np.uint8: + if m.max() <= 1.0: + m = (m.astype(np.float32) * 255.0).astype(np.uint8) + else: + m = m.astype(np.uint8) + # Many mask PNGs encode binary mask with alpha 0/255; use >0 to be robust + m = (m > 0).astype(np.uint8) * 255 + return m + + # numpy array + if isinstance(mask_data, np.ndarray): + return _ensure_size_u8(mask_data) + # PIL image + if isinstance(mask_data, Image.Image): + return _ensure_size_u8(np.array(mask_data.convert("L"))) + # Replicate FileOutput (URL-like) + if ReplicateFileOutput is not None and isinstance(mask_data, ReplicateFileOutput): + try: + url = getattr(mask_data, "url", None) + if isinstance(url, str) and url.startswith("http"): + m = download_mask(url) + return _ensure_size_u8(m) + except Exception as e: + print(f"[Segment] Failed to read Replicate FileOutput: {e}") + return None + # string forms + if isinstance(mask_data, str): + s = mask_data.strip() + if s.startswith("http://") or s.startswith("https://"): + try: + m = download_mask(s) # already 0..255 grayscale + return _ensure_size_u8(m) + except Exception as e: + print(f"[Segment] Failed to download mask: {e}") + return None + if s.startswith("data:image"): + try: + _, b64 = s.split(",", 1) + img_bytes = base64.b64decode(b64) + img = Image.open(BytesIO(img_bytes)) + # Prefer alpha channel if present + if "A" in img.getbands(): + m = np.array(img.getchannel("A")) + else: + m = np.array(img.convert("L")) + return _ensure_size_u8(m) + except Exception as e: + print(f"[Segment] Failed to decode data URI mask: {e}") + return None + # Fallback: some providers return raw base64-encoded PNG without data URI prefix + try: + img_bytes = base64.b64decode(s) + img = Image.open(BytesIO(img_bytes)) + if "A" in img.getbands(): + m = np.array(img.getchannel("A")) + else: + m = np.array(img.convert("L")) + return _ensure_size_u8(m) + except Exception: + pass + print("[Segment] Unknown mask string format; skipping") + return None + # dict forms + if isinstance(mask_data, dict): + # quick recursive search for any url/data string + def _find_any_image_string(obj): + try: + if isinstance(obj, str) and (obj.startswith("http") or obj.startswith("data:image")): + return obj + if isinstance(obj, dict): + for vv in obj.values(): + s = _find_any_image_string(vv) + if s: + return s + if isinstance(obj, (list, tuple)): + for vv in obj: + s = _find_any_image_string(vv) + if s: + return s + except Exception: + pass + return None + + # Handle COCO RLE formats (uncompressed only) + def _decode_coco_rle(rle_obj): + try: + if isinstance(rle_obj, dict) and isinstance(rle_obj.get("counts"), list) and "size" in rle_obj: + counts = rle_obj["counts"] + Hh, Ww = rle_obj["size"] + flat = np.zeros(Hh * Ww, dtype=np.uint8) + idx = 0 + val = 0 + for c in counts: + end = idx + int(c) + flat[idx:end] = val + idx = end + val = 255 - val + return flat.reshape((Hh, Ww)) + # Compressed RLE (string counts) not supported without pycocotools + return None + except Exception as e: + print(f"[Segment] Failed to decode RLE: {e}") + return None + + # Top-level RLE + if "rle" in mask_data and isinstance(mask_data["rle"], (dict,)): + decoded = _decode_coco_rle(mask_data["rle"]) + if decoded is not None: + return _ensure_size_u8(decoded) + # Some schemas put counts/size at top-level + if "counts" in mask_data and "size" in mask_data: + decoded = _decode_coco_rle({"counts": mask_data["counts"], "size": mask_data["size"]}) + if decoded is not None: + return _ensure_size_u8(decoded) + + # Try common fields carrying URLs or data URIs + for k in ["mask", "url", "image", "overlay", "png", "combined_mask"]: + v = mask_data.get(k) + if isinstance(v, str): + return _coerce_mask_to_numpy(v, target_hw) + if isinstance(v, dict): + # nested dict possibly with url or data + for kk in ["url", "image", "overlay", "data", "png"]: + vv = v.get(kk) + if isinstance(vv, str): + return _coerce_mask_to_numpy(vv, target_hw) + # nested numeric array + if isinstance(vv, (list, tuple)): + arr = np.array(vv) + if arr.ndim >= 2: + return _ensure_size_u8(arr) + # If dict has numeric array directly under known keys + for k in ["data", "array", "segmentation", "mask_array"]: + v = mask_data.get(k) + if isinstance(v, (list, tuple)): + arr = np.array(v) + if arr.ndim >= 2: + return _ensure_size_u8(arr) + # If dict has a single string value somewhere, try it + for v in mask_data.values(): + if isinstance(v, str): + return _coerce_mask_to_numpy(v, target_hw) + # Final attempt: recursively search any nested url or data-uri string + s_any = _find_any_image_string(mask_data) + if s_any: + return _coerce_mask_to_numpy(s_any, target_hw) + print("[Segment] Unknown dict mask format; skipping") + return None + # list-of-lists (numeric mask) + if isinstance(mask_data, (list, tuple)): + try: + arr = np.array(mask_data) + if arr.ndim >= 2: + return _ensure_size_u8(arr) + except Exception: + pass + return None + # unknown + return None + + +def segment( + image: Image.Image, + detection_results: List[Any], + segmentator, + processor=None, + device=None, + polygon_refinement=False, +): + """ + Segment objects using cloud-based Replicate SAM. + OPTIMIZED: More aggressive downscaling and fallback strategies. + """ + W, H = image.size + + # OPTIMIZATION: More aggressive downscaling + max_side_img = 768 # Reduced from 1024 + scale_img = 1.0 + image_for_sam = image + if max(W, H) > max_side_img: + scale_img = max_side_img / float(max(W, H)) + new_size = (int(round(W * scale_img)), int(round(H * scale_img))) + image_for_sam = image.resize(new_size, Image.BILINEAR) + print(f"[Segment] Using downscaled image for SAM: {(W, H)} -> {new_size}") + + # Run SAM once with timeout handling + buf_img = BytesIO() + image_for_sam.save(buf_img, format="PNG") + buf_img.seek(0) + + try: + output = segmentator(buf_img) + parsed = _parse_sam_output(output) + print(f"[Segment] global SAM outputs={len(parsed)}") + + # Coerce all masks to original size (H, W) + masks_np: List[np.ndarray] = [] + for j, md in enumerate(parsed): + m = _coerce_mask_to_numpy(md, target_hw=(H, W)) + if m is not None: + try: + nz = int((m > 0).sum()) + print(f"[Segment] mask[{j}] nz={nz}") + except Exception: + pass + masks_np.append(m) + + except Exception as e: + print(f"[Segment] SAM failed or timed out: {e}") + print("[Segment] Falling back to box-fill masks for all detections") + masks_np = [] + + results_with_masks = [] + # Assign best-overlap mask to each detection + for idx, det in enumerate(detection_results): + box = det.box + xmin, ymin, xmax, ymax = map(int, [box.xmin, box.ymin, box.xmax, box.ymax]) + xmin = max(0, min(xmin, W - 1)) + xmax = max(0, min(xmax, W)) + ymin = max(0, min(ymin, H - 1)) + ymax = max(0, min(ymax, H)) + if xmax <= xmin or ymax <= ymin: + print(f"[Segment] skip invalid box at idx {idx}: {(xmin, ymin, xmax, ymax)}") + results_with_masks.append(det) + continue + + box_area = max(1, (xmax - xmin) * (ymax - ymin)) + best_idx = -1 + best_iou = 0.0 + best_metrics = None + + # Only try mask matching if we have masks + if masks_np: + for k, m in enumerate(masks_np): + mask_area = int((m > 0).sum()) + # Skip masks that are basically full-frame (likely combined mask) + if mask_area / float(W * H) > 0.8: + continue + sub = m[ymin:ymax, xmin:xmax] + overlap = int((sub > 0).sum()) + if overlap == 0: + continue + # IoU with the detection box region + iou = overlap / float(mask_area + box_area - overlap + 1e-6) + if iou > best_iou: + best_iou = iou + best_idx = k + best_metrics = (overlap, mask_area) + + if best_idx >= 0 and best_iou > 0: + det.mask = masks_np[best_idx] + results_with_masks.append(det) + if idx < 5: + ov, ma = best_metrics if best_metrics else (0, 0) + print(f"[Segment] attach mask {best_idx} to det {idx}, overlap={ov}, mask_area={ma}, box_area={box_area}, iou={best_iou:.4f}") + else: + print(f"[Segment] no suitable mask for det {idx} — box fill fallback (box_area={box_area})") + full_mask = np.zeros((H, W), dtype=np.uint8) + full_mask[ymin:ymax, xmin:xmax] = 255 + det.mask = full_mask + results_with_masks.append(det) + + return results_with_masks + +# ------------------------------ +# Device helper +# ------------------------------ +def get_device(): + if torch.cuda.is_available(): + device = torch.device("cuda") + print("CUDA is available. Using GPU.") + elif torch.backends.mps.is_available(): + device = torch.device("mps") + print("MPS is available! Using Apple GPU.") + else: + device = torch.device("cpu") + print("Using CPU.") + return device diff --git a/ArtKrit/script/composition/server.py b/ArtKrit/script/composition/server.py new file mode 100644 index 0000000..bfbc975 --- /dev/null +++ b/ArtKrit/script/composition/server.py @@ -0,0 +1,187 @@ +# import os +# import time +# from flask import Flask, request, jsonify +# from flask_cors import CORS +# from run_models import * +# from composition_utils import fit_lines, line_leftmost_to_rightmost +# import cv2 + +# app = Flask(__name__) +# CORS(app, resources={r"/*": {"origins": "*"}}) + +# # Initialize models at server startup +# detector, segmentator, processor = init_models() + +# @app.route("/process_image", methods=['POST']) +# def process_image(): +# print("Processing image...") +# json_data = request.get_json() +# print(json_data) +# t0 = time.time() + +# ## read in the image file from the request and save it to a temporary file and get the path +# image_path = json_data["file_path"] + +# labels = [l.strip() for l in json_data["text_prompt"].split(",") if l.strip()] # Clean labels +# threshold_bbox = 0.3 +# polygon_refinement = True + +# if (not labels) and (len(json_data["custom_rectangles"]) == 0): +# return jsonify({"error": "No labels or custom rectangles provided"}) + +# if isinstance(image_path, str): +# image = load_image(image_path) + +# t_load = time.time() +# print(f"[Server] Calling Replicate GroundingDINO (labels={labels}, threshold={threshold_bbox})") +# detections = detect(image, labels, detector, threshold_bbox) +# t_detect = time.time() +# print(f"[Server] GroundingDINO call finished in {t_detect - t_load:.2f}s") +# for custom_rectangle in json_data["custom_rectangles"]: +# detections.append(DetectionResult.from_dict( +# { +# "score": 1.0, +# "label": "custom_rectangle", +# "box": { +# "xmin": int(custom_rectangle[0]), +# "ymin": int(custom_rectangle[1]), +# "xmax": int(custom_rectangle[2]), +# "ymax": int(custom_rectangle[3]) +# } +# } +# )) + +# # Cloud-only: do not request local device; SAM runs on Replicate +# print("[Server] Calling Replicate SAM model for segmentation") +# detections = segment(image, detections, segmentator, processor, None, polygon_refinement) +# t_segment = time.time() +# print(f"[Server] Replicate SAM call finished in {t_segment - t_detect:.2f}s") + +# image_array = np.array(image) + +# visualizaion_parameters = { +# # "polygon_epsilon": 0.008, +# "polygon_epsilon": json_data["polygon_epsilon"] * 1e-3, +# "point_radius": 1e-2, +# "line_fit_tol": 0.04, +# "line_radius": 1e-1 +# } + +# annotated_image, ploygon_contours_list, lines_list, points_to_draw = annotate(image_array,detections, visualizaion_parameters) +# t_annotate = time.time() + +# # Plot lines_list on the annotated_image +# top_lines = min(len(lines_list), 1) +# for j in range(top_lines): +# p1, p2 = lines_list[j] +# cv2.line(annotated_image, tuple(p1), tuple(p2), (0, 255, 0), 20) +# annotated_image_lines = cv2.cvtColor(annotated_image, cv2.COLOR_BGR2RGB) +# cv2.imwrite("temp/krita_temp_detection_res.png", img=annotated_image) + +# # Timing summary +# try: +# print(f"[Timing] load={t_load - t0:.2f}s detect={t_detect - t_load:.2f}s segment={t_segment - t_detect:.2f}s annotate={t_annotate - t_segment:.2f}s total={t_annotate - t0:.2f}s") +# except Exception as e: +# print(f"[Timing] error computing timings: {e}") + +# result = {"ploygon_contours": ploygon_contours_list, "composition_lines": lines_list, "points": points_to_draw} +# return jsonify(result) + +# @app.route('/regenerate_lines', methods=['POST']) +# def regenerate_lines(): +# """ +# Regenerate composition lines from manually adjusted points +# without calling the detection models again +# """ +# try: +# json_data = request.get_json() +# points = json_data.get('points', []) +# polygon_contours = json_data.get('polygon_contours', []) + +# if not points: +# return jsonify({'error': 'Points are required'}), 400 + +# if not polygon_contours: +# return jsonify({'error': 'Polygon contours are required'}), 400 + +# print(f"[Server] Regenerating lines from {len(points)} manually adjusted points") +# t0 = time.time() + +# # Convert points to the format expected by fit_lines +# # Each point needs to be associated with a polygon index +# points_with_index = assign_points_to_polygons(points, polygon_contours) + +# # Create a dummy image array for shape information +# # We need to infer image dimensions from the polygon contours +# max_x = max(max(p[0] for p in contour) for contour in polygon_contours) +# max_y = max(max(p[1] for p in contour) for contour in polygon_contours) +# image_shape = (int(max_y) + 1, int(max_x) + 1, 3) +# dummy_image = np.zeros(image_shape, dtype=np.uint8) + +# # Use the line fitting logic WITHOUT calling annotate or sample_contour_points +# line_fit_tol = 0.04 +# inlier_threshold = 0.05 +# lines = fit_lines(points_with_index, dummy_image, line_fit_tol=line_fit_tol, inlier_threshold=inlier_threshold) + +# t_generate = time.time() +# print(f"[Server] Generated {len(lines)} composition lines in {t_generate - t0:.2f}s") + +# # Convert lines to the same format as process_image endpoint +# lines_list = [] +# for line in lines: +# p1, p2 = line_leftmost_to_rightmost(line) +# lines_list.append([[int(p1[0]), int(p1[1])], [int(p2[0]), int(p2[1])]]) + +# response = { +# 'composition_lines': lines_list, +# 'num_points': len(points), +# 'num_lines': len(lines_list) +# } + +# return jsonify(response) + +# except Exception as e: +# print(f"[Server] Error regenerating lines: {str(e)}") +# import traceback +# traceback.print_exc() +# return jsonify({'error': str(e)}), 500 + +# def assign_points_to_polygons(points, polygon_contours): +# """ +# Assign each point to the polygon it's closest to. + +# Args: +# points: List of [x, y] coordinates +# polygon_contours: List of polygon contours (each is a list of [x, y] coordinates) + +# Returns: +# List of tuples: [(point, polygon_index), ...] +# """ +# points_with_index = [] + +# for point in points: +# point_array = np.array(point, dtype=np.float32) +# min_distance = float('inf') +# closest_polygon_idx = 0 + +# # Find the closest polygon to this point +# for poly_idx, contour in enumerate(polygon_contours): +# contour_array = np.array(contour, dtype=np.float32).reshape(-1, 2) + +# # Calculate distance to each point in the contour +# distances = np.linalg.norm(contour_array - point_array, axis=1) +# min_dist_to_contour = np.min(distances) + +# if min_dist_to_contour < min_distance: +# min_distance = min_dist_to_contour +# closest_polygon_idx = poly_idx + +# points_with_index.append((point, closest_polygon_idx)) + +# print(f"[Server] Assigned {len(points)} points to {len(polygon_contours)} polygons") +# return points_with_index + + +# if __name__ == '__main__': +# print("Starting server") +# app.run(host='localhost', port=5001, debug=True) \ No newline at end of file diff --git a/ArtKrit/script/value_color/category_data.py b/ArtKrit/script/value_color/category_data.py new file mode 100644 index 0000000..7d52f6e --- /dev/null +++ b/ArtKrit/script/value_color/category_data.py @@ -0,0 +1,254 @@ +from abc import ABC, abstractmethod +from typing import Dict, List, Tuple, Any, Optional +import cv2 +import numpy as np + +class BlobInfo: + """Holds information about a detected blob: pixel points, bounding box, and contours.""" + def __init__(self): + self.points: List[Tuple[int, int]] = [] + self.bbox: Optional[Tuple[int, int, int, int]] = None + self.contours: List[np.ndarray] = [] + +class BaseCategoryData(ABC): + """ + Abstract base class for a category of analysis (value vs. color). + Stores all state (maps, blobs, dominants, matches) and defines the + generic blob-creation routine. Subclasses implement how to + threshold and extract dominants. + """ + def __init__(self): + # pixel maps: hex_code -> list of (x, y) + self.canvas_map: Dict[str, List[Tuple[int, int]]] = {} + self.reference_map: Dict[str, List[Tuple[int, int]]] = {} + # dominant features: list of (feature, hex_code) + self.canvas_dominant: List[Tuple[Any, str]] = [] + self.reference_dominant: List[Tuple[Any, str]] = [] + # blob info: hex_code -> BlobInfo + self.canvas_blobs: Dict[str, BlobInfo] = {} + self.reference_blobs: Dict[str, BlobInfo] = {} + # matched pairs: canvas_hex -> reference_hex + self.matched_pairs: Dict[str, str] = {} + + @abstractmethod + def threshold_mask(self, image: np.ndarray, feature: Any) -> np.ndarray: + """ + Given an image and a feature (value or RGB tuple), + return a binary mask where that feature "occurs". + """ + pass + + @abstractmethod + def extract_dominant(self, image: np.ndarray, **kwargs) -> List[Tuple[Any, str]]: + """ + Perform clustering on `image` and return a list of + (feature, hex_code) tuples. Feature is either a + grayscale value or an RGB tuple. + """ + pass + + def create_map_with_blobs(self, + image: np.ndarray, + use_canvas: bool = True) -> None: + """ + Generic blob-creation routine: + - Chooses canvas vs. reference lists and maps + - For each dominant feature, builds a mask via threshold_mask + - Finds contours, records points and bbox + """ + dominants = self.canvas_dominant if use_canvas else self.reference_dominant + data_map = self.canvas_map if use_canvas else self.reference_map + data_blobs = self.canvas_blobs if use_canvas else self.reference_blobs + + data_map.clear() + data_blobs.clear() + + for feature, hex_code in dominants: + # Threshold the image to get a binary mask + mask = self.threshold_mask(image, feature) + # Find contours on the mask + contours, _ = cv2.findContours(mask.copy(), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + blob = BlobInfo() + blob.contours = contours + + # Flatten all contour points + all_pts: List[Tuple[int, int]] = [] + for cnt in contours: + for pt in cnt.reshape(-1, 2): + x, y = int(pt[0]), int(pt[1]) + all_pts.append((x, y)) + blob.points = all_pts + + # Compute bounding box + if all_pts: + xs, ys = zip(*all_pts) + x_min, x_max = min(xs), max(xs) + y_min, y_max = min(ys), max(ys) + blob.bbox = (x_min, y_min, x_max - x_min, y_max - y_min) + + data_map[hex_code] = blob.points + data_blobs[hex_code] = blob + +class ValueData(BaseCategoryData): + """Concrete for grayscale 'value' analysis.""" + def threshold_mask(self, image, value_level): + """ + Build a binary mask where pixels fall within ±15 levels + of the specified grayscale value. + """ + # Compute lower/upper bounds, clipped to valid [0,255] range + lower = max(0, value_level - 15) + upper = min(255, value_level + 15) + + # Ensure it's a single‐channel grayscale image + # If it's already 2D, use it directly; otherwise convert from BGR. + if image.ndim == 2: + gray = image + else: + gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) + + # Use OpenCV’s inRange to build the mask: + # pixels in [lower, upper] → 255 (white), others → 0 (black) + mask = cv2.inRange(gray, lower, upper) + + # Return the resulting binary mask + return mask + + def extract_dominant(self, image, num_values=5, **kwargs): + """ + Find the most frequent gray levels in the image via k-means clustering. + Returns a list of (value, hex_code) tuples, sorted by frequency descending. + """ + # Convert to a single-channel grayscale array if needed + # If the image is already 2D, assume it's grayscale. Otherwise convert. + if image.ndim == 2: + gray = image + else: + gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) + + # Flatten to a long column of pixel-values for clustering + # From (H, W) → (H*W, 1) and cast to float32 + pixels = gray.reshape(-1, 1).astype(np.float32) + + # Decide on number of clusters, can’t exceed unique levels or num_values + unique_levels = np.unique(pixels).size + K = min(num_values, unique_levels) + + # Define k-means stopping criteria: max 10 iters or epsilon 1.0 + criteria = (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 10, 1.0) + + # Run k-means + # attempts=3 for robustness, center-init via KMEANS_PP_CENTERS + _, labels, centers = cv2.kmeans( + pixels, + K, + None, + criteria, + attempts=3, + flags=cv2.KMEANS_PP_CENTERS + ) + + # Count how many pixels fell into each cluster + counts = np.bincount(labels.flatten(), minlength=K) + + # Sort clusters by descending size (most common first) + order = np.argsort(-counts) + + # Build output list, converting each center to int + hex string + result = [] + for idx in order: + # Center is a 1-element array, take its first entry + value = int(centers[idx][0]) + # Format as hex triplet (R=G=B=value) + hex_code = f"#{value:02x}{value:02x}{value:02x}" + result.append((value, hex_code)) + + return result + + +class ColorData(BaseCategoryData): + """Concrete for RGB 'color' analysis.""" + def threshold_mask(self, image, rgb): + """ + Create a binary mask highlighting pixels whose color + lies within ±18 of the target RGB tuple. + """ + # Build lower/upper bounds for each channel, clamped to [0,255] + lower = np.array([ + max(0, rgb[0] - 18), + max(0, rgb[1] - 18), + max(0, rgb[2] - 18), + ], dtype=np.uint8) + upper = np.array([ + min(255, rgb[0] + 18), + min(255, rgb[1] + 18), + min(255, rgb[2] + 18), + ], dtype=np.uint8) + + # Make sure the image is RGB (3 channels). + # If it has 4 channels (e.g. BGRA), convert it. + if image.shape[2] == 3: + img_rgb = image + else: + img_rgb = cv2.cvtColor(image, cv2.COLOR_BGRA2RGB) + + # Use OpenCV to threshold: pixels in [lower,upper] become 255, rest 0 + mask = cv2.inRange(img_rgb, lower, upper) + + # Return the binary mask + return mask + + def extract_dominant(self, image, num_values=6, **kwargs): + """ + Find the most frequent colors in the image using k-means clustering. + Returns a list of ((r, g, b), hex_code) tuples, sorted by frequency descending. + """ + # Ensure it is an RGB array (drop alpha if present) + # If the image has 3 channels, assume it's already RGB + if image.shape[2] == 3: + # rgb = image + rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) + else: + # Convert BGRA → RGB, discarding alpha + rgb = cv2.cvtColor(image, cv2.COLOR_BGRA2RGB) + + # Reshape to a list of pixels for clustering + # From (H, W, 3) → (H*W, 3), and cast to float32 + pixels = rgb.reshape(-1, 3).astype(np.float32) + + # Decide how many clusters to find + # We can’t ask for more clusters than we have pixels + K = min(num_values, pixels.shape[0]) + + # Set up k-means stopping criteria + # Stop either after 10 iterations or when movement < 1.0 + criteria = (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 10, 1.0) + + # Run k-means clustering + # flags=cv2.KMEANS_PP_CENTERS for smart seeding + _, labels, centers = cv2.kmeans( + pixels, + K, + None, + criteria, + attempts=3, + flags=cv2.KMEANS_PP_CENTERS + ) + + # Count how many pixels fall into each cluster + counts = np.bincount(labels.flatten(), minlength=K) + + # Sort cluster indices by descending population + order = np.argsort(-counts) + + # Build the result list, converting centers → ints + hex codes + result = [] + for idx in order: + # Cluster center is floating-point BGR; convert to ints and reorder if needed + r, g, b = [int(c) for c in centers[idx]] + # Format as hex string + hex_code = f"#{r:02x}{g:02x}{b:02x}" + result.append(((r, g, b), hex_code)) + + return result + diff --git a/ArtKrit/script/value_color/helpers/color_conversion.py b/ArtKrit/script/value_color/helpers/color_conversion.py new file mode 100644 index 0000000..a135749 --- /dev/null +++ b/ArtKrit/script/value_color/helpers/color_conversion.py @@ -0,0 +1,69 @@ +"""Utlity functions for color conversion""" + +import numpy as np +import cv2 + +def hex_to_lab(hex_color): + # 1) Hex → sRGB [0–1] + r = int(hex_color[1:3], 16) / 255.0 + g = int(hex_color[3:5], 16) / 255.0 + b = int(hex_color[5:7], 16) / 255.0 + + def inv_gamma(c): + return ((c + 0.055) / 1.055) ** 2.4 if c > 0.04045 else c / 12.92 + r_lin, g_lin, b_lin = inv_gamma(r), inv_gamma(g), inv_gamma(b) + + # 2) Linear RGB → XYZ (D65) + x = (r_lin * 0.4124564 + g_lin * 0.3575761 + b_lin * 0.1804375) * 100 + y = (r_lin * 0.2126729 + g_lin * 0.7151522 + b_lin * 0.0721750) * 100 + z = (r_lin * 0.0193339 + g_lin * 0.1191920 + b_lin * 0.9503041) * 100 + + # 3) Normalize by D65 white point + x, y, z = x / 95.047, y / 100.000, z / 108.883 + + # 4) XYZ → Lab + def f(t): + return t ** (1/3) if t > 0.008856 else (7.787 * t + 16/116) + fx, fy, fz = f(x), f(y), f(z) + L = 116 * fy - 16 + a = 500 * (fx - fy) + b = 200 * (fy - fz) + return L, a, b + +def rgb_to_hsv(rgb): + """ + Convert RGB tuple to HSV tuple using OpenCV, + converting to standard display ranges. + + Args: + rgb (tuple): RGB color values (0-255 range) + + Returns: + tuple: HSV color values in standard ranges + Hue: 0-360 + Saturation: 0-100 + Value: 0-100 + """ + print(f"DEBUG: Input RGB = {rgb}") + # Create a single pixel numpy array in BGR format (OpenCV's native format) + bgr_pixel = np.uint8([[[rgb[2], rgb[1], rgb[0]]]]) # ← BGR order: B, G, R + print(f"DEBUG: BGR pixel = {bgr_pixel[0, 0]}") # ← Add this line + + # Convert BGR to HSV + hsv_pixel = cv2.cvtColor(bgr_pixel, cv2.COLOR_BGR2HSV) + + # Extract OpenCV HSV values + h, s, v = hsv_pixel[0, 0] + + # Convert to standard ranges + # Hue: 0-179 -> 0-360 + h_standard = int(h) * 2 + print(f"hue: {h}") + + # Saturation: 0-255 -> 0-100 + s_standard = round(s / 255 * 100) + + # Value: 0-255 -> 0-100 + v_standard = round(v / 255 * 100) + + return (int(h_standard), int(s_standard), int(v_standard)) \ No newline at end of file diff --git a/ArtKrit/script/value_color/helpers/color_separation_tool.py b/ArtKrit/script/value_color/helpers/color_separation_tool.py new file mode 100644 index 0000000..87ca19e --- /dev/null +++ b/ArtKrit/script/value_color/helpers/color_separation_tool.py @@ -0,0 +1,430 @@ +from krita import Krita, InfoObject +import cv2 +import numpy as np +from PyQt5.QtCore import Qt, QPoint, pyqtSignal +from PyQt5.QtWidgets import (QLabel, QScrollArea, QVBoxLayout, QWidget, + QHBoxLayout, QPushButton, QColorDialog, QSlider) +from PyQt5.QtGui import QImage, QPixmap, QColor +import os +import sys +from PIL import Image +import json +from datetime import datetime + + +class ClusterHoverLabel(QLabel): + clusterHovered = pyqtSignal(int) + + def __init__(self, parent=None): + super().__init__(parent) + self.setMouseTracking(True) + self.setAlignment(Qt.AlignCenter) + self.original_img = None + self.cluster_labels = None + self.dominant_colors = None + self.cluster_groups = None + self.current_pixmap = None + self._pixmap_offset = QPoint(0, 0) + self._pixmap_scale = 1.0 + self.background_color = (255, 255, 255) # Default white + + + def setBackgroundColor(self, color): + """Set the background color for highlighting (RGB tuple)""" + self.background_color = color + + def setImageData(self, original_img, labels, colors, groups): + self.original_img = original_img + self.cluster_labels = labels + self.dominant_colors = colors + self.cluster_groups = groups + + h, w = original_img.shape[:2] + bytes_per_line = 3 * w + qimg = QImage(original_img.data, w, h, bytes_per_line, QImage.Format_RGB888) + self.current_pixmap = QPixmap.fromImage(qimg) + self.setPixmap(self.scalePixmap(self.current_pixmap)) + + def scalePixmap(self, pixmap): + if pixmap.isNull(): + return pixmap + + # Get available size safely + try: + parent_widget = self.parent() + if parent_widget and hasattr(parent_widget, 'width') and hasattr(parent_widget, 'height'): + available_width = max(100, parent_widget.width() - 20) + available_height = parent_widget.height() - 20 + else: + # Fallback values if parent is not available + available_width = 400 + available_height = 300 + except RuntimeError: + # Parent widget has been deleted, use fallback values + available_width = 400 + available_height = 300 + + # Calculate scaled size maintaining aspect ratio + pixmap_ratio = pixmap.width() / pixmap.height() + available_ratio = available_width / available_height + + if pixmap_ratio > available_ratio: + # Width is the limiting factor + scaled_width = available_width + scaled_height = int(scaled_width / pixmap_ratio) + else: + # Height is the limiting factor + scaled_height = available_height + scaled_width = int(scaled_height * pixmap_ratio) + + # Store scaling factor and offset for accurate coordinate mapping + self._pixmap_scale = scaled_width / pixmap.width() + + # Calculate the centered position + self._pixmap_offset = QPoint( + (self.width() - scaled_width) // 2 if hasattr(self, 'width') else 0, + (self.height() - scaled_height) // 2 if hasattr(self, 'height') else 0 + ) + + return pixmap.scaled( + scaled_width, scaled_height, + Qt.KeepAspectRatio, + Qt.SmoothTransformation + ) + + def mouseMoveEvent(self, event): + try: + if self.original_img is None or self.current_pixmap is None: + return + + # Convert mouse position to pixmap coordinates + pos = event.pos() + + # Get current pixmap dimensions + if not self.pixmap() or self.pixmap().isNull(): + return + + pixmap_width = self.pixmap().width() + pixmap_height = self.pixmap().height() + + # Check if mouse is within the pixmap area + if not (self._pixmap_offset.x() <= pos.x() < self._pixmap_offset.x() + pixmap_width and + self._pixmap_offset.y() <= pos.y() < self._pixmap_offset.y() + pixmap_height): + self.leaveEvent(event) + return + + # Calculate position in original image coordinates + pixmap_x = pos.x() - self._pixmap_offset.x() + pixmap_y = pos.y() - self._pixmap_offset.y() + + orig_x = int(pixmap_x / self._pixmap_scale) + orig_y = int(pixmap_y / self._pixmap_scale) + + # Ensure position is within bounds + if (0 <= orig_x < self.cluster_labels.shape[1] and + 0 <= orig_y < self.cluster_labels.shape[0]): + cluster_idx = self.cluster_labels[orig_y, orig_x] + group_idx = self.cluster_groups[cluster_idx] + + # Highlight all clusters in this group with custom background color + highlighted = np.full_like(self.original_img, self.background_color, dtype=np.uint8) + mask = np.isin(self.cluster_labels, [c for c, g in enumerate(self.cluster_groups) if g == group_idx]) + highlighted[mask] = self.original_img[mask] + + h, w = highlighted.shape[:2] + bytes_per_line = 3 * w + qimg = QImage(highlighted.data, w, h, bytes_per_line, QImage.Format_RGB888) + self.current_pixmap = QPixmap.fromImage(qimg) + self.setPixmap(self.scalePixmap(self.current_pixmap)) + + self.clusterHovered.emit(group_idx) + except RuntimeError: + # Widget has been deleted, ignore the event + pass + + def leaveEvent(self, event): + if self.original_img is not None: + h, w = self.original_img.shape[:2] + bytes_per_line = 3 * w + qimg = QImage(self.original_img.data, w, h, bytes_per_line, QImage.Format_RGB888) + self.current_pixmap = QPixmap.fromImage(qimg) + self.setPixmap(self.scalePixmap(self.current_pixmap)) + super().leaveEvent(event) + + def resizeEvent(self, event): + if self.current_pixmap: + self.setPixmap(self.scalePixmap(self.current_pixmap)) + super().resizeEvent(event) + + +class ColorSeparationTool: + def __init__(self, parent): + self.parent = parent + self.current_image = None + self.current_labels = None + self.current_colors = None + self.current_groups = None + self.image_label = None + self.background_color = QColor(255, 255, 255) # Default white + self.bg_color_button = None + self.color_groups_slider = None + self.color_groups_label = None + self.num_color_groups = 8 # Default number of color groups + + + def get_json_path(self): + """Get the path to the logs JSON file""" + home_dir = os.path.expanduser("~") + logs_folder = os.path.join(home_dir, "ArtKrit_logs") + os.makedirs(logs_folder, exist_ok=True) + return os.path.join(logs_folder, "logs.json") + + def append_log_entry(self, action, message): + """Append a log entry to the JSON log file""" + self.save_png_on_button_press(action) + json_path = self.get_json_path() + + try: + with open(json_path, "r") as f: + data = json.load(f) + except FileNotFoundError: + data = {"logs": []} + + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + + data["logs"].append({ + "timestamp": timestamp, + "action": action, + "message": message + }) + + with open(json_path, "w") as f: + json.dump(data, f, indent=4) + + print(f"Logged: {action}") + + + def save_png_on_button_press(self, action): + """Save the current document as PNG when a button is pressed""" + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + safe_action = action.replace(" ", "_") + + home_dir = os.path.expanduser("~") + base_folder = os.path.join(home_dir, "ArtKrit_logs") + images_folder = os.path.join(base_folder, "canvas_images") + os.makedirs(images_folder, exist_ok=True) + + doc = Krita.instance().activeDocument() + if doc is not None: + current_path = doc.fileName() + if not current_path: + print("Please save your document first.") + return + + png_path = os.path.join(images_folder, f"{timestamp}_{safe_action}.png") + + doc.setBatchmode(True) + options = InfoObject() + options.setProperty('compression', 5) + options.setProperty('alpha', True) + doc.exportImage(png_path, options) + doc.setBatchmode(False) + + + def cleanup(self): + """Clean up resources to prevent memory leaks""" + if self.image_label: + try: + self.image_label.deleteLater() + except: + pass + self.image_label = None + + if self.bg_color_button: + try: + self.bg_color_button.deleteLater() + except: + pass + self.bg_color_button = None + + if self.color_groups_slider: + try: + self.color_groups_slider.deleteLater() + except: + pass + self.color_groups_slider = None + + if self.color_groups_label: + try: + self.color_groups_label.deleteLater() + except: + pass + self.color_groups_label = None + + self.current_image = None + self.current_labels = None + self.current_colors = None + self.current_groups = None + + def create_color_separation_ui(self): + """Create the UI components for color separation""" + # Main container + main_container = QWidget(self.parent.color_tab) + main_layout = QVBoxLayout(main_container) + + # Color picker controls + controls_layout = QHBoxLayout() + + # Background color button + self.bg_color_button = QPushButton("Background Color") + self.bg_color_button.clicked.connect(self.choose_background_color) + self.update_color_button_style() + controls_layout.addWidget(self.bg_color_button) + + controls_layout.addStretch() + main_layout.addLayout(controls_layout) + + # Color groups slider + slider_layout = QHBoxLayout() + self.color_groups_label = QLabel(f"Color Groups: {self.num_color_groups}") + slider_layout.addWidget(self.color_groups_label) + + self.color_groups_slider = QSlider(Qt.Horizontal) + self.color_groups_slider.setMinimum(3) + self.color_groups_slider.setMaximum(30) + self.color_groups_slider.setValue(self.num_color_groups) + self.color_groups_slider.setTickPosition(QSlider.TicksBelow) + self.color_groups_slider.setTickInterval(3) + self.color_groups_slider.valueChanged.connect(self.on_slider_changed) + slider_layout.addWidget(self.color_groups_slider) + + main_layout.addLayout(slider_layout) + + # Scrollable image container + scroll_area = QScrollArea() + scroll_area.setWidgetResizable(True) + scroll_area.setHorizontalScrollBarPolicy(Qt.ScrollBarAsNeeded) + scroll_area.setVerticalScrollBarPolicy(Qt.ScrollBarAsNeeded) + + img_cont = QWidget() + v = QVBoxLayout(img_cont) + self.image_label = ClusterHoverLabel(img_cont) + self.image_label.setMinimumSize(200, 200) + self.image_label.clusterHovered.connect(self.update_cluster_info) + v.addWidget(self.image_label) + + scroll_area.setWidget(img_cont) + main_layout.addWidget(scroll_area) + + return main_container, self.image_label + + def on_slider_changed(self, value): + """Handle slider value change""" + self.num_color_groups = value + if self.color_groups_label: + self.color_groups_label.setText(f"Color Groups: {value}") + + # Reprocess the image with new number of color groups + self.update_cluster_count() + self.append_log_entry("Color Groups Changed", f"New number of color groups: {value}") + + def choose_background_color(self): + """Open color picker dialog to choose background color""" + color = QColorDialog.getColor(self.background_color, self.parent, "Choose Background Color") + + if color.isValid(): + self.background_color = color + self.update_color_button_style() + + # Update the image label's background color + if self.image_label: + self.image_label.setBackgroundColor((color.red(), color.green(), color.blue())) + self.append_log_entry("Background Color Changed", f"New color: {color.name()}") + + + def update_color_button_style(self): + """Update the button style to show the current background color""" + if self.bg_color_button: + r, g, b = self.background_color.red(), self.background_color.green(), self.background_color.blue() + # Calculate contrast color for text (black or white) + brightness = (r * 299 + g * 587 + b * 114) / 1000 + text_color = "black" if brightness > 128 else "white" + + self.bg_color_button.setStyleSheet( + f"background-color: rgb({r}, {g}, {b}); " + f"color: {text_color}; " + f"border: 1px solid #888; " + f"padding: 5px 10px; " + f"border-radius: 3px;" + ) + + def process_reference_image(self, color_reference_image): + """Process the stored reference image for color analysis""" + if color_reference_image is not None: + # Convert color space appropriately + if len(color_reference_image.shape) == 2: # Grayscale + self.current_image = cv2.cvtColor(color_reference_image, cv2.COLOR_GRAY2RGB) + elif color_reference_image.shape[2] == 4: # RGBA + self.current_image = cv2.cvtColor(color_reference_image, cv2.COLOR_BGRA2RGB) + else: # BGR + self.current_image = cv2.cvtColor(color_reference_image, cv2.COLOR_BGR2RGB) + + self.update_cluster_count() + + def update_cluster_count(self): + """Recompute dominant color clusters and update the image label accordingly.""" + if self.current_image is None: + return + + # Check if image_label exists before using it + if self.image_label is None: + return + + self.parent.color_data.reference_dominant = self.parent.color_data.extract_dominant( + self.current_image, + num_values=self.num_color_groups + ) + dominant_colors = self.parent.color_data.reference_dominant + + # Convert dominant colors to the format expected by the rest of the code + h, w = self.current_image.shape[:2] + self.current_labels = np.zeros((h, w), dtype=np.int32) + self.current_colors = np.array([rgb for (rgb, _) in dominant_colors], dtype=np.uint8) + + # Create a mask for each color and assign labels + self.current_groups = [] + dominant_colors_array = np.array([rgb for rgb, _ in dominant_colors]) + reshaped_image = self.current_image.reshape((-1, 3)) + distances = np.linalg.norm(reshaped_image[:, np.newaxis] - dominant_colors_array, axis=2) + closest_color_indices = np.argmin(distances, axis=1) + self.current_labels = closest_color_indices.reshape(self.current_image.shape[:2]) + self.current_groups = np.unique(closest_color_indices).tolist() + + # Update display + unique_groups = len(set(self.current_groups)) + + # Send to display - with error handling + try: + self.image_label.setImageData( + self.current_image, + self.current_labels, + self.current_colors, + self.current_groups + ) + except RuntimeError: + # Image label has been deleted, recreate it + pass + + def update_cluster_info(self, group_idx): + if self.current_colors is None or self.current_groups is None: + return + + # Count pixels in this group + group_mask = np.isin(self.current_groups, [group_idx]) + group_colors = self.current_colors[group_mask] + + # Display group info + color_info = " | ".join( + f"RGB({c[0]}, {c[1]}, {c[2]})" + for c in group_colors + ) \ No newline at end of file diff --git a/ArtKrit/script/value_color/helpers/image_conversion.py b/ArtKrit/script/value_color/helpers/image_conversion.py new file mode 100644 index 0000000..8f86021 --- /dev/null +++ b/ArtKrit/script/value_color/helpers/image_conversion.py @@ -0,0 +1,72 @@ +"""Utlity functions for image conversion""" +import cv2 + +def _to_grayscale(img): + """ + Convert any image format to grayscale. + Handles: BGR, BGRA, RGB, RGBA, or already grayscale + Returns: Single channel grayscale image + """ + if img is None: + return None + + # Already grayscale + if len(img.shape) == 2: + return img + + # Has color channels + channels = img.shape[2] + + if channels == 4: # RGBA or BGRA + # Krita uses BGRA, cv2.imread with alpha uses BGRA + return cv2.cvtColor(img, cv2.COLOR_BGRA2GRAY) + elif channels == 3: # RGB or BGR + # cv2.imread without alpha uses BGR + return cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) + + return img + +def _to_bgr(img): + """ + Convert any image format to BGR (3 channels). + Handles: Grayscale, BGRA, already BGR + Returns: 3-channel BGR image + """ + if img is None: + return None + + # Already BGR + if len(img.shape) == 3 and img.shape[2] == 3: + return img + + # Grayscale - convert to BGR + if len(img.shape) == 2: + return cv2.cvtColor(img, cv2.COLOR_GRAY2BGR) + + # BGRA - remove alpha channel + if len(img.shape) == 3 and img.shape[2] == 4: + return cv2.cvtColor(img, cv2.COLOR_BGRA2BGR) + + return img + +def _to_rgb_for_display(img): + """ + Convert any image format to RGB for Qt display. + Handles: Grayscale, BGR, BGRA, already RGB + Returns: 3-channel RGB image ready for QImage + """ + if img is None: + return None + + # Grayscale - convert to RGB + if len(img.shape) == 2: + return cv2.cvtColor(img, cv2.COLOR_GRAY2RGB) + + channels = img.shape[2] + + if channels == 4: # BGRA + return cv2.cvtColor(img, cv2.COLOR_BGRA2RGB) + elif channels == 3: # BGR + return cv2.cvtColor(img, cv2.COLOR_BGR2RGB) + + return img \ No newline at end of file diff --git a/ArtKrit/script/value_color/helpers/lasso_fill_tool.py b/ArtKrit/script/value_color/helpers/lasso_fill_tool.py new file mode 100644 index 0000000..d22905b --- /dev/null +++ b/ArtKrit/script/value_color/helpers/lasso_fill_tool.py @@ -0,0 +1,297 @@ +from krita import Krita, ManagedColor, InfoObject +from PyQt5.QtCore import Qt, QTimer +from PyQt5.QtWidgets import QGroupBox, QVBoxLayout, QPushButton +from PyQt5.QtGui import QColor +import numpy as np +import os +import sys +from PIL import Image +import json +from datetime import datetime + + +class LassoFillTool: + def __init__(self, parent): + self.parent = parent + self.currentFillColor = None + self.selectionTimer = QTimer() + self.selectionTimer.setSingleShot(True) + self.selectionTimer.timeout.connect(self.checkSelection) + + + def get_json_path(self): + """Get the path to the logs JSON file""" + home_dir = os.path.expanduser("~") + logs_folder = os.path.join(home_dir, "ArtKrit_logs") + os.makedirs(logs_folder, exist_ok=True) + return os.path.join(logs_folder, "logs.json") + + def append_log_entry(self, action, message): + """Append a log entry to the JSON log file""" + self.save_png_on_button_press(action) + json_path = self.get_json_path() + + try: + with open(json_path, "r") as f: + data = json.load(f) + except FileNotFoundError: + data = {"logs": []} + + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + + data["logs"].append({ + "timestamp": timestamp, + "action": action, + "message": message + }) + + with open(json_path, "w") as f: + json.dump(data, f, indent=4) + + print(f"Logged: {action}") + + def save_png_on_button_press(self, action): + """Save the current document as PNG when a button is pressed""" + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + safe_action = action.replace(" ", "_") + + home_dir = os.path.expanduser("~") + base_folder = os.path.join(home_dir, "ArtKrit_logs") + images_folder = os.path.join(base_folder, "canvas_images") + os.makedirs(images_folder, exist_ok=True) + + doc = Krita.instance().activeDocument() + if doc is not None: + current_path = doc.fileName() + if not current_path: + print("Please save your document first.") + return + + png_path = os.path.join(images_folder, f"{timestamp}_{safe_action}.png") + + doc.setBatchmode(True) + options = InfoObject() + options.setProperty('compression', 5) + options.setProperty('alpha', True) + doc.exportImage(png_path, options) + doc.setBatchmode(False) + + # In lasso_fill_tool.py + def create_fill_widgets(self): + """Create the fill options UI components""" + fill_grp = QGroupBox("Fill Options") + fv = QVBoxLayout(fill_grp) + + # Create the fill color button + self.fillColorButton = QPushButton("Select Fill Color") + self.fillColorButton.clicked.connect(self.selectFillColor) # Connect to lasso tool method + self.fillColorButton.setStyleSheet("background-color: #ffffff;") + + # Create the fill button + self.fillButton = QPushButton("Fill Selection") + self.fillButton.clicked.connect(self.fillSelection) # Connect to lasso tool method + self.fillButton.setEnabled(False) + + fv.addWidget(self.fillColorButton) + fv.addWidget(self.fillButton) + fill_grp.setLayout(fv) + fill_grp.setVisible(False) + + return fill_grp, self.fillColorButton, self.fillButton + + def activateLassoTool(self): + """Activate Krita's lasso selection tool and prepare the fill options UI.""" + krita_instance = Krita.instance() + + action = krita_instance.action('KisToolSelectContiguous') + if action: + action.trigger() + self.parent.lassoButton.setStyleSheet("background-color: #AED6F1;") + + self.parent.fillGroup.setVisible(True) + + self.selectionTimer.start(500) + + QTimer.singleShot(500, lambda: self.parent.lassoButton.setStyleSheet("")) + self.append_log_entry("lasso tool", "Lasso tool activated") + + def selectFillColor(self): + """Open the HS picker seeded by the current selection's average value.""" + # Extract the dominant value from the selection + doc = Krita.instance().activeDocument() + if doc: + selection = doc.selection() + if selection: + node = doc.activeNode() + average_value = self.extractAverageValueFromSelection(node, selection) + + if average_value is not None: + # Open the color picker with the extracted value + from ..value_color import CustomHSColorPickerDialog + dialog = CustomHSColorPickerDialog(self.parent, average_value) + dialog.exec_() + + selectedColor = dialog.selectedColor() + if selectedColor.isValid(): + self.currentFillColor = selectedColor + self.fillColorButton.setStyleSheet(f"background-color: {selectedColor.name()};") + else: + # Show error message or use default value + print("Warning: Could not extract value from selection. Using default value.") + self.parent.color_feedback_label.setText("⚠️ Could not extract value from selection. Please ensure you have a valid selection with visible pixels.") + + # Use a reasonable default value (e.g., mid-tone) + from ..value_color import CustomHSColorPickerDialog + dialog = CustomHSColorPickerDialog(self.parent, 128) + dialog.exec_() + + selectedColor = dialog.selectedColor() + if selectedColor.isValid(): + self.currentFillColor = selectedColor + self.fillColorButton.setStyleSheet(f"background-color: {selectedColor.name()};") + self.append_log_entry("lasso fill color", f"Selected lasso fill color: {selectedColor.name()}") + + + def fillSelection(self): + """Fill the current selection with the selected fill color.""" + krita_instance = Krita.instance() + doc = krita_instance.activeDocument() + selection = doc.selection() + node = doc.activeNode() + average_value = self.extractAverageValueFromSelection(node, selection) + if average_value is None: + print("Failed to extract average value from selection") + return + + # Get the selected H and S from the color picker + selected_hue = self.currentFillColor.hue() + selected_saturation = self.currentFillColor.saturation() + + # Create the new color with the extracted value and selected H and S + new_color = QColor.fromHsv(selected_hue, selected_saturation, average_value) + print(f"New fill color: {new_color.name()}") + + # Convert QColor to ManagedColor for Krita + managedColor = ManagedColor("RGBA", "U8", "") + managedColor.setComponents([ + new_color.blueF(), + new_color.greenF(), + new_color.redF(), + 1.0 # Fully opaque + ]) + + # Set the foreground color + if krita_instance.activeWindow() and krita_instance.activeWindow().activeView(): + view = krita_instance.activeWindow().activeView() + view.setForeGroundColor(managedColor) + print("Foreground color set to new color") + + # Trigger the fill tool action + fillToolAction = krita_instance.action('KritaFill/KisToolFill') + if fillToolAction: + print("Triggering fill tool action") + fillToolAction.trigger() + + QTimer.singleShot(100, lambda: self.triggerFillForeground(krita_instance)) + self.append_log_entry("lasso fill initiated", f"Initiated lasso fill with color: {new_color.name()}") + + else: + print("Could not find fill tool action") + + + + def triggerFillForeground(self, krita_instance): + """Triggers the fill_foreground action after the fill tool is activated.""" + fillAction = krita_instance.action('fill_foreground') + if fillAction: + fillAction.trigger() + + + def extractAverageValueFromSelection(self, node, selection): + """ + Extracts the dominant brightness (value) from the selected area using the HSV color space. + Returns None if no valid value can be extracted. + """ + try: + print("Extracting pixel data from selection...") + + # Check if selection is valid + if not selection or selection.width() == 0 or selection.height() == 0: + print("Invalid or empty selection") + return None + + # Get the pixel data from the selected area + pixel_data = node.projectionPixelData( + selection.x(), selection.y(), selection.width(), selection.height() + ).data() + + if not pixel_data or len(pixel_data) == 0: + print("No pixel data available") + return None + + pixels = [] + for i in range(0, len(pixel_data), 4): + # Skip if we don't have enough data + if i + 3 >= len(pixel_data): + break + + r = pixel_data[i] + g = pixel_data[i + 1] + b = pixel_data[i + 2] + # Only add pixels that aren't completely transparent (alpha > 0) + a = pixel_data[i + 3] + if a > 0: # Only consider non-transparent pixels + pixels.append((r, g, b)) + + if not pixels: + print("No non-transparent pixels in selection") + return None + + # Calculate the frequency of each brightness (value) level + value_counts = {} + for r, g, b in pixels: + # Convert RGB to HSV + color = QColor(r, g, b) + if color.isValid(): + h, s, v, _ = color.getHsv() + + # Count the frequency of each value + if v in value_counts: + value_counts[v] += 1 + else: + value_counts[v] = 1 + + if not value_counts: + print("No valid colors found in selection") + return None + + # Find the dominant value (the one with the highest frequency) + dominant_value = max(value_counts, key=value_counts.get) + print(f"Dominant value (brightness): {dominant_value}") + + # Ensure the value is within valid range (0-255) + if 0 <= dominant_value <= 255: + return dominant_value + else: + print(f"Invalid value extracted: {dominant_value}") + return None + + except Exception as e: + print(f"Error extracting dominant value: {str(e)}") + return None + + def checkSelection(self): + doc = Krita.instance().activeDocument() + if doc: + selection = doc.selection() + if selection: + self.fillButton.setEnabled(True) + else: + self.fillButton.setEnabled(False) + + if self.parent.fillGroup.isVisible(): + self.selectionTimer.start(500) + + + + \ No newline at end of file diff --git a/ArtKrit/script/value_color/helpers/matching_algo.py b/ArtKrit/script/value_color/helpers/matching_algo.py new file mode 100644 index 0000000..a848cf5 --- /dev/null +++ b/ArtKrit/script/value_color/helpers/matching_algo.py @@ -0,0 +1,59 @@ +"""Utility functions for value and color matching algorithms""" + +from . import color_conversion +import numpy as np + +def calculate_bbox_overlap(bbox1, bbox2): + """Calculate the overlap ratio between two bounding boxes.""" + if not bbox1 or not bbox2: + return 0.0 + + x1, y1, w1, h1 = bbox1 + x2, y2, w2, h2 = bbox2 + + # Calculate coordinates of the intersection + x_left = max(x1, x2) + y_top = max(y1, y2) + x_right = min(x1 + w1, x2 + w2) + y_bottom = min(y1 + h1, y2 + h2) + + if x_right < x_left or y_bottom < y_top: + # No overlap + return 0.0 + + intersection_area = (x_right - x_left) * (y_bottom - y_top) + + bbox1_area = w1 * h1 + bbox2_area = w2 * h2 + + union_area = bbox1_area + bbox2_area - intersection_area + + # Calculate IoU + iou = intersection_area / union_area if union_area > 0 else 0.0 + + return iou + +def calculate_color_similarity(hex1, hex2, is_color_analysis=False): + """ + Calculate value‐similarity S_val between two hex colors per: + S_val = 1 - (1/3) * ||ΔLab|| after normalizing L*∈[0,1], a*∈[0,1], b*∈[0,1]. + """ + # Convert both hex colors to Lab + L1, a1, b1 = color_conversion.hex_to_lab(hex1) + L2, a2, b2 = color_conversion.hex_to_lab(hex2) + + Ln1, an1, bn1 = L1 / 100.0, (a1 + 128) / 255.0, (b1 + 128) / 255.0 + Ln2, an2, bn2 = L2 / 100.0, (a2 + 128) / 255.0, (b2 + 128) / 255.0 + + # Euclidean distance in the normalized cube (max = √3) + delta = np.sqrt( + (Ln1 - Ln2)**2 + + (an1 - an2)**2 + + (bn1 - bn2)**2 + ) + + # Formula: S_val = 1 – (1/3) * Δ + similarity = 1.0 - (delta / 3.0) + + # Clamp to [0,1] + return max(0.0, min(1.0, similarity)) \ No newline at end of file diff --git a/ArtKrit/script/value_color/helpers/text_feedback.py b/ArtKrit/script/value_color/helpers/text_feedback.py new file mode 100644 index 0000000..269f993 --- /dev/null +++ b/ArtKrit/script/value_color/helpers/text_feedback.py @@ -0,0 +1,114 @@ +"""Utility functions for generating natural language feedback""" + +def compute_hue_feedback(canvas_hsv, reference_hsv, hue_ranges): + """Given two HSV triples (hue, saturation, value), return a string + explaining how your canvas hue/saturation compares to the reference.""" + # Unpack the HSV components for clarity + canvas_hue, canvas_saturation, canvas_value = canvas_hsv + ref_hue, ref_saturation, ref_value = reference_hsv + + # Helper: map a hue angle to its descriptive range name + def _hue_label(angle): + for start, end, name in hue_ranges: + if start <= angle < end: + return name + return "unknown" + + # Find which named bucket each hue falls into + canvas_label = _hue_label(canvas_hue) + ref_label = _hue_label(ref_hue) + + hue_diff = canvas_hue - ref_hue + if ref_saturation: + sat_diff = (canvas_saturation - ref_saturation) / ref_saturation * 100 + else: + sat_diff = 0 + + if canvas_label == ref_label: + if hue_diff == 0: + hue_feedback = "Your canvas matches the reference exactly in hue." + else: + direction = "warmer" if hue_diff > 0 else "cooler" + intensity = "slightly" if abs(hue_diff) < 10 else "quite" + hue_feedback = ( + f"You're in the same {canvas_label} range as the reference, " + f"but the reference is {intensity} {direction}. " + ) + else: + direction = "warmer" if hue_diff > 0 else "cooler" + hue_feedback = ( + f"Your canvas sits in the {canvas_label} range, " + f"while the reference sits in the {ref_label} range. " + f"which is {direction}." + ) + + if sat_diff > 5: + sat_feedback = f"Your color is {abs(sat_diff):.1f}% more saturated than the reference." + elif sat_diff < -5: + sat_feedback = f"Your color is {abs(sat_diff):.1f}% less saturated than the reference." + else: + sat_feedback = "Your color has similar saturation to the reference." + + value_feedback = get_value_feedback(canvas_value, ref_value) + + feedback = "HSV Differences:\n" + feedback += f"Hue Difference: {hue_diff}°\n" + feedback += hue_feedback + "\n" + feedback += sat_feedback + "\n" + feedback += f"Value Difference: {value_feedback}\n" + + return feedback + +def get_color_feedback(canvas_hsv, reference_hsv, hue_ranges): + """Generate feedback about the color comparison.""" + feedback = compute_hue_feedback(canvas_hsv, reference_hsv, hue_ranges) + return feedback + +def get_value_feedback(canvas_value, ref_value): + """Generate feedback about the value comparison.""" + # Calculate the difference between canvas and reference values + canvas_value = int(canvas_value) + ref_value = int(ref_value) + if ref_value == 0: + if canvas_value == 0: + return "Both the canvas and reference are pure black (value = 0)." + else: + return ("The reference is pure black (value = 0), " + "but the canvas has brightness. Canvas is lighter.") + + value_diff = (canvas_value - ref_value) / ref_value * 100 + feedback = "" + + # Set thresholds for differences + minor_threshold = 5 + moderate_threshold = 10 + major_threshold = 20 + + + if abs(value_diff) <= minor_threshold: + feedback = f"These values are closely matched (difference: {abs(value_diff):.1f}%)" + else: + if value_diff < 0: + # Canvas is darker than reference + if abs(value_diff) > major_threshold: + feedback = f"Canvas values are significantly too dark (it's {abs(value_diff):.1f}% darker than the reference.)" + elif abs(value_diff) > moderate_threshold: + feedback = f"Canvas values are moderately too dark (it's {abs(value_diff):.1f}% darker than the reference.)" + else: + feedback = f"Canvas values are slightly too dark (by {abs(value_diff):.1f}% darker than the reference.)" + else: + # Canvas is lighter than reference + if abs(value_diff) > major_threshold: + feedback = f"Canvas values are significantly too light (by {abs(value_diff):.1f}% lighter than the reference.)" + elif abs(value_diff) > moderate_threshold: + feedback = f"Canvas values are moderately too light (by {abs(value_diff):.1f}% lighter than the reference.)" + else: + feedback = f"Canvas values are slightly too light (by {abs(value_diff):.1f}% lighter than the reference.)" + + # if abs(value_diff) > minor_threshold: + # if value_diff < 0: + # feedback += "\nSuggestion: Try lightening this area to better match the reference." + # else: + # feedback += "\nSuggestion: Try darkening this area to better match the reference." + + return feedback \ No newline at end of file diff --git a/ArtKrit/script/value_color/value_color.py b/ArtKrit/script/value_color/value_color.py new file mode 100644 index 0000000..35c9765 --- /dev/null +++ b/ArtKrit/script/value_color/value_color.py @@ -0,0 +1,1633 @@ +"""Main user interface for value and color tabs""" + +from krita import DockWidget, DockWidgetFactory, DockWidgetFactoryBase, Krita, ManagedColor, InfoObject +from PyQt5.QtCore import Qt, QPoint, QTimer, pyqtSignal +from PyQt5.QtWidgets import ( + QWidget, QVBoxLayout, QSlider, QLabel, QPushButton, QButtonGroup, QRadioButton, + QHBoxLayout, QSplitter, QScrollArea, QDialog, QGroupBox, QSizePolicy +) +from PyQt5.QtGui import QImage, QPixmap, QPainter, QColor, QPainterPath, QConicalGradient, QBrush +import os +import sys +sys.path.append(os.path.expanduser("~/ddraw/lib/python3.10/site-packages")) +import cv2 +import numpy as np +from PIL import Image +import math +from sklearn.cluster import KMeans +from scipy.cluster.hierarchy import linkage, fcluster +import json +from datetime import datetime +from .category_data import ValueData, ColorData +from .helpers import color_conversion, image_conversion, matching_algo, text_feedback +from .helpers.lasso_fill_tool import LassoFillTool +from .helpers.color_separation_tool import ColorSeparationTool + +class ValueButton(QPushButton): + def __init__(self, value, hex_code, is_reference=False, parent=None): + super().__init__(parent) + self.value = value + self.hex_code = hex_code + self.is_reference = is_reference + self.matched_button = None + self.setMinimumSize(60, 60) + self.setMaximumSize(60, 60) + border_style = "2px solid #00FF00" if is_reference else "1px solid #888888" + self.setStyleSheet(f"background-color: {hex_code}; border: {border_style};") + self.setToolTip(hex_code) + + def set_matched_button(self, button): + self.matched_button = button + +class ValuePairWidget(QWidget): + clicked = pyqtSignal(object) # emitted when this pair is clicked + + def __init__(self, canvas_rgb, canvas_hex, ref_rgb, ref_hex, parent=None): + super().__init__(parent) + self.canvas_hex = canvas_hex + self.ref_hex = ref_hex + layout = QHBoxLayout() + layout.setSpacing(5) + self.setLayout(layout) + # Create buttons + self.canvas_button = ValueButton(canvas_rgb, canvas_hex, is_reference=False) + self.ref_button = ValueButton(ref_rgb, ref_hex, is_reference=True) + + # Connect both buttons to emit clicked(self) + self.canvas_button.clicked.connect(self._emit_clicked) + self.ref_button.clicked.connect(self._emit_clicked) + + # Create labels + canvas_label = QLabel(f"{canvas_rgb}") + canvas_label.setAlignment(Qt.AlignCenter) + canvas_label.setStyleSheet("color: white; background-color: #333333; padding: 2px;") + + ref_label = QLabel(f"{ref_rgb}") + ref_label.setAlignment(Qt.AlignCenter) + ref_label.setStyleSheet("color: white; background-color: #333333; padding: 2px;") + + arrow_label = QLabel("→") + arrow_label.setAlignment(Qt.AlignCenter) + arrow_label.setStyleSheet("color: #FFFF00; font-size: 16px; font-weight: bold;") + + layout.addWidget(self.canvas_button) + layout.addWidget(canvas_label) + layout.addWidget(arrow_label) + layout.addWidget(ref_label) + layout.addWidget(self.ref_button) + + def _emit_clicked(self): + """Emit signal when either button is clicked.""" + self.clicked.emit(self) + + def set_highlight(self, enabled): + """Highlight this pair green if selected.""" + if enabled: + self.setStyleSheet("background-color: rgba(0,255,0,60); border: 2px solid #00FF00;") + else: + self.setStyleSheet("") + + +class ValueColor(QWidget): + """Widget for loading images, applying filters, and showing value/color analyses.""" + def __init__(self, parent=None): + super().__init__(parent) + self.value_image = None + self.color_image = None + self.value_reference_image = None + self.color_reference_image = None + + # Separate canvas images for color and value + self.value_canvas_image = None # Grayscale canvas for value analysis + self.color_canvas_image = None # Color canvas for color analysis + + self.current_filter = None + + self.value_data = ValueData() + self.color_data = ColorData() + + self.value_pair_widgets = [] + self.color_pair_widgets = [] + + # Initialize lasso fill tool + self.lasso_fill_tool = LassoFillTool(self) + self.color_separation_tool = ColorSeparationTool(self) + + self.selectionTimer = QTimer() + self.selectionTimer.setSingleShot(True) + self.selectionTimer.timeout.connect(self.lasso_fill_tool.checkSelection) + + def cleanup(self): + """Clean up resources""" + if hasattr(self, 'color_separation_tool'): + self.color_separation_tool.cleanup() + if hasattr(self, 'lasso_fill_tool'): + pass + + def __del__(self): + self.cleanup() + + def export_pixmap(self, pixmap, action): + """Export a QPixmap as PNG to the logs directory""" + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + safe_action = action.replace(" ", "_") + + home_dir = os.path.expanduser("~") + export_folder = os.path.join(home_dir, "ArtKrit_logs", "artkrit_output_images") + os.makedirs(export_folder, exist_ok=True) + + filename = os.path.join(export_folder, f"{timestamp}_{safe_action}.png") + pixmap.save(filename, "PNG") + print(f"Saved filtered image as {filename}") + + + def export_filtered_image_as_png(self, filtered_image, action): + """Export a filtered image as PNG to the logs directory""" + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + safe_action = action.replace(" ", "_") + + home_dir = os.path.expanduser("~") + export_folder = os.path.join(home_dir, "ArtKrit_logs", "artkrit_output_images") + os.makedirs(export_folder, exist_ok=True) + + img = Image.fromarray(filtered_image) + filename = os.path.join(export_folder, f"{timestamp}_{safe_action}.png") + img.save(filename) + print(f"Saved filtered image as {filename}") + + def save_png_on_button_press(self, action): + """Save the current document as PNG when a button is pressed""" + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + safe_action = action.replace(" ", "_") + + home_dir = os.path.expanduser("~") + base_folder = os.path.join(home_dir, "ArtKrit_logs") + images_folder = os.path.join(base_folder, "canvas_images") + os.makedirs(images_folder, exist_ok=True) + + doc = Krita.instance().activeDocument() + if doc is not None: + current_path = doc.fileName() + if not current_path: + print("Please save your document first.") + return + + png_path = os.path.join(images_folder, f"{timestamp}_{safe_action}.png") + + doc.setBatchmode(True) + options = InfoObject() + options.setProperty('compression', 5) + options.setProperty('alpha', True) + doc.exportImage(png_path, options) + doc.setBatchmode(False) + + def get_json_path(self): + """Get the path to the logs JSON file""" + home_dir = os.path.expanduser("~") + logs_folder = os.path.join(home_dir, "ArtKrit_logs") + os.makedirs(logs_folder, exist_ok=True) + return os.path.join(logs_folder, "logs.json") + + def append_log_entry(self, action, message): + """Append a log entry to the JSON log file""" + self.save_png_on_button_press(action) + json_path = self.get_json_path() + + try: + with open(json_path, "r") as f: + data = json.load(f) + except FileNotFoundError: + data = {"logs": []} + + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + + data["logs"].append({ + "timestamp": timestamp, + "action": action, + "message": message + }) + + with open(json_path, "w") as f: + json.dump(data, f, indent=4) + + print(f"Logged: {action}") + + + def _make_preview_section(self, prefix, left_text, right_text): + """ + Creates: + - self.{prefix}_preview_container + - QHBoxLayout on it + - self.{prefix}_left_preview_label + - self.{prefix}_right_preview_label + - self.{prefix}_preview_splitter + Returns the container widget (so caller can add it to layout). + """ + container = QWidget() + layout = QHBoxLayout(container) + + left = QLabel(left_text) + left.setAlignment(Qt.AlignCenter) + left.setStyleSheet("background-color: black; color: white;") + left.setMinimumHeight(300) + setattr(self, f"{prefix}_left_preview_label", left) + + right = QLabel(right_text) + right.setAlignment(Qt.AlignCenter) + right.setStyleSheet("background-color: black; color: white;") + right.setMinimumHeight(300) + setattr(self, f"{prefix}_right_preview_label", right) + + splitter = QSplitter(Qt.Horizontal) + splitter.addWidget(left) + splitter.addWidget(right) + splitter.setSizes([200,200]) + setattr(self, f"{prefix}_preview_splitter", splitter) + + layout.addWidget(splitter) + setattr(self, f"{prefix}_preview_container", container) + + return container + + def _make_pairs_section(self, prefix, header_text): + """ + Creates: + - self.{prefix}_pairs_layout (QVBoxLayout) + - self.{prefix}_matched_pairs_label (QLabel header) + - self.{prefix}_pairs_container (+ its layout) + - QScrollArea containing that container. + Returns the QLayout (so caller can add it to the tab). + """ + pairs_layout = QVBoxLayout() + setattr(self, f"{prefix}_pairs_layout", pairs_layout) + + header = QLabel(header_text) + header.setStyleSheet("font-weight: bold; font-size: 14px;") + setattr(self, f"{prefix}_matched_pairs_label", header) + pairs_layout.addWidget(header) + + container = QWidget() + container_layout = QVBoxLayout(container) + setattr(self, f"{prefix}_pairs_container", container) + setattr(self, f"{prefix}_pairs_container_layout", container_layout) + + scroll = QScrollArea() + scroll.setWidgetResizable(True) + scroll.setWidget(container) + scroll.setMinimumHeight(200) + scroll.setMaximumHeight(400) + pairs_layout.addWidget(scroll) + + return pairs_layout + + def _make_feedback_label(self, prefix, text): + """Creates a scrollable feedback label with consistent stylesheet""" + scroll = QScrollArea() + scroll.setWidgetResizable(True) + scroll.setMaximumHeight(150) + scroll.setMinimumHeight(80) + scroll.setHorizontalScrollBarPolicy(Qt.ScrollBarAlwaysOff) + scroll.setVerticalScrollBarPolicy(Qt.ScrollBarAsNeeded) + + lbl = QLabel(text) + lbl.setAlignment(Qt.AlignLeft | Qt.AlignTop) + lbl.setWordWrap(True) + lbl.setStyleSheet(""" + QLabel { + color: white; + background-color: #333333; + padding: 2px; + border-radius: 2px; + margin: 2px; + } + """) + + scroll.setWidget(lbl) + setattr(self, f"{prefix}_feedback_label", lbl) + + return scroll + + # Individual tab-creators + def create_value_tab(self): + """Create the tab for value analysis""" + self.value_tab = QWidget() + layout = QVBoxLayout(self.value_tab) + layout.setAlignment(Qt.AlignTop) + + # canvas picker + btn = QPushButton("Set Current Canvas") + btn.clicked.connect(self.show_current_canvas) + layout.addWidget(btn) + + # filter radios + slider + filter_group = QGroupBox("Filter Options") + h = QHBoxLayout(filter_group) + self.filter_group = QButtonGroup(filter_group) + for name in ("Gaussian","Bilateral","Median"): + rb = QRadioButton(name) + # re-create the old attributes so upload_image() still works: + if name == "Gaussian": + self.gaussian_radio = rb + elif name == "Bilateral": + self.bilateral_radio = rb + else: + self.median_radio = rb + + rb.clicked.connect(lambda _, n=name.lower(): self.filter_selected(n)) + self.filter_group.addButton(rb) + h.addWidget(rb) + layout.addWidget(filter_group) + + self.slider_label = QLabel("Kernel Size (%): 1.5%") + self.slider = QSlider(Qt.Horizontal) + self.slider.setRange(15,49) + self.slider.setValue(15) + self.slider.valueChanged.connect(self.update_kernel_size_label) + self.slider.valueChanged.connect(self.update_preview) + self.slider_label.hide() + self.slider.hide() + layout.addWidget(self.slider_label) + layout.addWidget(self.slider) + + # --- SHARED SECTIONS --- + layout.addWidget(self._make_preview_section( + "value", "Canvas", "Reference" + )) + fb_btn = QPushButton("Get Value Feedback") + fb_btn.clicked.connect(self.get_feedback_value) + self.value_feedback_btn = fb_btn + layout.addWidget(fb_btn) + + layout.addLayout(self._make_pairs_section( + "value", "Value Pairs (Canvas → Reference):" + )) + + layout.addWidget(self._make_feedback_label( + "value", "Click 'Get Value Feedback' to analyze the canvas values" + )) + + def create_color_tab(self): + """Create the tab for color analysis""" + self.color_tab = QWidget() + layout = QVBoxLayout(self.color_tab) + layout.setAlignment(Qt.AlignTop) + self.setWindowTitle("Color Cluster Matcher") + + # Process & zoom controls + proc = QPushButton("Process Reference Image") + proc.clicked.connect(self.process_reference_image) + layout.addWidget(proc) + + zoom_h = QHBoxLayout() + for txt, slot in (("- Zoom out", self.zoom_out),("+ Zoom in", self.zoom_in)): + btn = QPushButton(txt) + btn.clicked.connect(slot) + zoom_h.addWidget(btn) + layout.addLayout(zoom_h) + + # Button to pop out/dock color separation tool + self.color_sep_toggle_btn = QPushButton("↗ Pop Out Color Separation") + self.color_sep_toggle_btn.clicked.connect(self.toggle_color_separation_window) + layout.addWidget(self.color_sep_toggle_btn) + + # Create color separation UI (initially embedded) + self.color_sep_container, self.image_label = self.color_separation_tool.create_color_separation_ui() + self.color_sep_parent_layout = layout # Store reference to parent layout + layout.addWidget(self.color_sep_container) + + # Initialize floating window as None + self.color_sep_floating_window = None + self.color_sep_is_floating = False + + # Color tools & fill options + tools = QGroupBox("Color Tools") + tv = QVBoxLayout(tools) + + # Create the color button + # self.colorButton = QPushButton("Select Color") + # self.colorButton.clicked.connect(self.selectColor) + # tv.addWidget(self.colorButton) + + # Create the lasso button and store it as an attribute + self.lassoButton = QPushButton("Lasso Fill Tool") + self.lassoButton.clicked.connect(self.lasso_fill_tool.activateLassoTool) + tv.addWidget(self.lassoButton) + + # Get fill widgets from lasso fill tool and add them + self.fillGroup, self.fillColorButton, self.fillButton = self.lasso_fill_tool.create_fill_widgets() + tv.addWidget(self.fillGroup) + + layout.addWidget(tools) + + # --- SHARED SECTIONS --- + layout.addWidget(self._make_preview_section( + "color", "Canvas", "Reference" + )) + cfbtn = QPushButton("Get Color Feedback") + cfbtn.clicked.connect(self.get_feedback_color) + self.color_feedback_btn = cfbtn + layout.addWidget(cfbtn) + + layout.addLayout(self._make_pairs_section( + "color", "Color Pairs (Canvas → Reference):" + )) + + layout.addWidget(self._make_feedback_label( + "color", "Process reference image first" + )) + + # Initialize image data storage + self.current_image = None + self.current_labels = None + self.current_colors = None + self.current_groups = None + + def toggle_color_separation_window(self): + """Toggle the color separation tool between embedded and floating window""" + # Toggle state + if self.color_sep_is_floating: + self.dock_color_separation() + self.append_log_entry("toggle color separation window close", "Toggled color separation tool window state: docked") + else: + self.pop_out_color_separation() + self.append_log_entry("toggle color separation window open", "Toggled color separation tool window state: popped out") + + def pop_out_color_separation(self): + """Pop out the color separation tool into a floating window""" + # Create floating window + self.color_sep_floating_window = QDialog(self) + self.color_sep_floating_window.setWindowTitle("Color Separation Tool") + self.color_sep_floating_window.resize(800, 600) + + # Create layout for the dialog + dialog_layout = QVBoxLayout(self.color_sep_floating_window) + + # Remove container from parent and add to dialog + self.color_sep_container.setParent(None) + dialog_layout.addWidget(self.color_sep_container) + + # Update button text + self.color_sep_toggle_btn.setText("↙ Dock Color Separation") + self.color_sep_is_floating = True + + # Show the window + self.color_sep_floating_window.show() + + # Connect close event to dock back + self.color_sep_floating_window.finished.connect(self.on_floating_window_closed) + + def dock_color_separation(self): + """Dock the color separation tool back into the color tab""" + # Remove from floating window + if self.color_sep_floating_window: + self.color_sep_container.setParent(None) + self.color_sep_floating_window.close() + self.color_sep_floating_window = None + + # Find the position to insert (after the toggle button) + toggle_index = self.color_sep_parent_layout.indexOf(self.color_sep_toggle_btn) + self.color_sep_parent_layout.insertWidget(toggle_index + 1, self.color_sep_container) + + # Update button text + self.color_sep_toggle_btn.setText("↗ Pop Out Color Separation") + self.color_sep_is_floating = False + + def on_floating_window_closed(self): + """Handle when the floating window is closed by the user""" + if self.color_sep_is_floating: + self.dock_color_separation() + self.append_log_entry("toggle color separation window close", "Toggled color separation tool window state: docked") + + def process_reference_image(self): + """Process the stored reference image for color analysis""" + if hasattr(self, 'color_reference_image') and self.color_reference_image is not None: + if hasattr(self, 'color_separation_tool') and self.color_separation_tool is not None: + self.color_separation_tool.process_reference_image(self.color_reference_image) + self.append_log_entry("process ref img for color", "Processed reference image for color separation") + + def update_cluster_count(self): + """Recompute dominant color clusters and update the image label accordingly.""" + if self.current_image is None: + return + + self.color_data.reference_dominant = self.color_data.extract_dominant( + self.current_image, + num_values=15 + ) + dominant_colors = self.color_data.reference_dominant + + # Convert dominant colors to the format expected by the rest of the code + h, w = self.current_image.shape[:2] + self.current_labels = np.zeros((h, w), dtype=np.int32) + self.current_colors = np.array([rgb for (rgb, _) in dominant_colors], dtype=np.uint8) + + # Create a mask for each color and assign labels + self.current_groups = [] + dominant_colors_array = np.array([rgb for rgb, _ in dominant_colors]) + reshaped_image = self.current_image.reshape((-1, 3)) + distances = np.linalg.norm(reshaped_image[:, np.newaxis] - dominant_colors_array, axis=2) + closest_color_indices = np.argmin(distances, axis=1) + self.current_labels = closest_color_indices.reshape(self.current_image.shape[:2]) + self.current_groups = np.unique(closest_color_indices).tolist() + + # Update display + unique_groups = len(set(self.current_groups)) + + # Send to display + self.image_label.setImageData( + self.current_image, + self.current_labels, + self.current_colors, + self.current_groups + ) + + def update_cluster_info(self, group_idx): + self.color_separation_tool.update_cluster_info(group_idx) + + def display_preview(self, img, is_color): + """Display the image in the preview label.""" + if img is None: + return + + # Convert to RGB for display using helper + rgb_img = image_conversion._to_rgb_for_display(img) + height, width, channels = rgb_img.shape + bytes_per_line = channels * width + qimage = QImage(rgb_img.data, width, height, bytes_per_line, QImage.Format_RGB888) + + pixmap = QPixmap.fromImage(qimage) + if is_color: + self.color_left_preview_label.setPixmap(pixmap.scaled( + self.color_left_preview_label.width(), + self.color_left_preview_label.height(), + Qt.KeepAspectRatio)) + else: + self.value_left_preview_label.setPixmap(pixmap.scaled( + self.value_left_preview_label.width(), + self.value_left_preview_label.height(), + Qt.KeepAspectRatio)) + + def upload_image(self, file_path): + """Load the image, initialize blank/reference canvases for color & value, and display them.""" + # Read image (preserving alpha channel if present) + self.color_image = cv2.imread(file_path, cv2.IMREAD_UNCHANGED) + self.color_reference_image = self.color_image.copy() + + # Create a blank canvas image of the same size + blank_canvas = np.zeros_like(self.color_image) + + # Display with blank canvas on left and reference on right + self.display_split_view(blank_canvas, self.color_image, True) + + # For value analysis - convert to grayscale + self.value_image = image_conversion._to_grayscale(cv2.imread(file_path, cv2.IMREAD_UNCHANGED)) + self.value_reference_image = self.value_image.copy() + + # Create a blank canvas for value view + blank_value_canvas = np.zeros_like(self.value_image) + + # Display with blank canvas on left and reference on right + self.display_split_view(blank_value_canvas, self.value_image, False) + + # Default to Gaussian filter if none selected + if not self.current_filter: + self.gaussian_radio.setChecked(True) + self.filter_selected("gaussian") + else: + self.update_preview() + + def filter_selected(self, filter_type): + """Set the active filter type, reveal the slider, and refresh the preview.""" + self.append_log_entry(f"applying {filter_type} filter", f"Selected filter: {filter_type}") + self.current_filter = filter_type + self.slider_label.show() + self.slider.show() + self.update_preview() + + def update_kernel_size_label(self, value): + """Reflect the slider's value as a percentage in the label text.""" + kernel_percentage = value / 10.0 + self.slider_label.setText(f"Kernel Size (%): {kernel_percentage:.1f}%") + + def update_preview(self): + """Apply the chosen filter to value and canvas images and update the split view.""" + if self.value_image is None or self.current_filter is None: + return + + # Make a copy of the image to work with + working_image = image_conversion._to_grayscale(self.value_image.copy()) + + # Calculate kernel size + kernel_percentage = self.slider.value() / 10.0 + hw_max = max(working_image.shape[:2]) + kernel_size = max(3, int(hw_max * (kernel_percentage / 100.0))) + if kernel_size % 2 == 0: + kernel_size += 1 + + # Apply filter to reference image + if self.current_filter == "gaussian": + self.filtered_image = cv2.GaussianBlur(working_image, (kernel_size, kernel_size), 0) + elif self.current_filter == "bilateral": + sigma_color = 75 + sigma_space = 75 + self.filtered_image = cv2.bilateralFilter(working_image, kernel_size, sigma_color, sigma_space) + elif self.current_filter == "median": + self.filtered_image = cv2.medianBlur(working_image, kernel_size) + + # Check if canvas_image has any non-zero values + if hasattr(self, 'value_canvas_image') and self.value_canvas_image is not None and np.any(self.value_canvas_image): + canvas_gray = image_conversion._to_grayscale(self.value_canvas_image.copy()) + # Apply the same filter to canvas image + if self.current_filter == "gaussian": + self.filtered_canvas = cv2.GaussianBlur(canvas_gray, (kernel_size, kernel_size), 0) + elif self.current_filter == "bilateral": + self.filtered_canvas = cv2.bilateralFilter(canvas_gray, kernel_size, sigma_color, sigma_space) + elif self.current_filter == "median": + self.filtered_canvas = cv2.medianBlur(canvas_gray, kernel_size) + else: + self.filtered_canvas = np.zeros_like(self.filtered_image) + + self.display_split_view(self.filtered_canvas, self.filtered_image, False) + self.export_filtered_image_as_png(self.filtered_image, f"applied_{self.current_filter}_filter_image, with {kernel_size} kernel") + self.export_filtered_image_as_png(self.filtered_canvas, f"applied_{self.current_filter}_filter_canvas, with {kernel_size} kernel") + self.append_log_entry("Update value preview reference", f"Applied {self.current_filter} filter with kernel size {kernel_size} to reference") + self.append_log_entry("Update value preview canvas", f"Applied {self.current_filter} filter with kernel size {kernel_size} to canvas") + + + + def get_feedback_value(self): + """Extract dominant values from the canvas and compare with the reference.""" + document = Krita.instance().activeDocument() + if not document: + self.value_feedback_label.setText("⚠️ No document is open") + return + + if self.value_reference_image is None: + self.value_feedback_label.setText("⚠️ Please upload a reference image first") + return + + # Apply default Gaussian filter if none selected + if not self.current_filter: + self.gaussian_radio.setChecked(True) + self.current_filter = "gaussian" + self.slider_label.show() + self.slider.show() + + if self.filtered_image is None: + self.update_preview() + + if self.filtered_image is not None: + # Extract the 20 most dominant values from reference + self.value_data.reference_dominant = ( + self.value_data.extract_dominant(self.filtered_image, num_values=20) + ) + self.value_data.create_map_with_blobs(self.filtered_image, use_canvas=False) + + # Get current canvas data + pixel_array = self.get_canvas_data() + if pixel_array is None: + return + + # Convert to grayscale using helper + pixel_array_gray = image_conversion._to_grayscale(pixel_array) + self.value_canvas_image = pixel_array_gray + + self.update_preview() + pixel_array_gray = self.filtered_canvas + + # Extract dominant values from canvas and reference + self.value_data.canvas_dominant = (self.value_data.extract_dominant(pixel_array_gray, num_values=5)) + + # Create value maps and blob information + self.value_data.create_map_with_blobs(pixel_array_gray, use_canvas=True) + + # Match canvas and reference values using spatial information + self.match_values(is_color_analysis=False) + + # Update UI with value pair widgets + self.update_pairs( + self.value_data, + self.value_pairs_container_layout, + lambda v: v, # no transform for raw grayscale levels + self.show_pair_regions_value + ) + + # Create and display the initial overview with all matched pairs + self.show_all_matched_pairs(False) + self.value_feedback_label.setText("✅ Found dominant values. Click on any value pair to see detailed comparison.") + self.append_log_entry("Get value feedback", "Requested value feedback analysis") + + def get_feedback_color(self): + """Extract dominant colors from the canvas and compare with the reference.""" + document = Krita.instance().activeDocument() + if not document: + self.color_feedback_label.setText("⚠️ No document is open") + return + + if self.color_reference_image is None: + self.color_feedback_label.setText("⚠️ Please upload a reference image first") + return + + # Get current canvas data + active_layer = document.activeNode() + doc_width, doc_height = document.width(), document.height() + pixel_data = active_layer.pixelData(0, 0, doc_width, doc_height) + pixel_array = np.frombuffer(pixel_data, dtype=np.uint8).reshape(doc_height, doc_width, -1) + + # Downsample pixel array to half size using cv2.resize + pixel_array = cv2.resize(pixel_array, (doc_width//2, doc_height//2), interpolation=cv2.INTER_AREA) + + # Store the original canvas image without format conversion + self.color_canvas_image = pixel_array.copy() + self.color_filtered_canvas = self.color_canvas_image.copy() + + # Create a version for analysis (RGB) + analysis_image = pixel_array.copy() + ref_analysis_image = self.color_reference_image.copy() + + # Extract the 15 most dominant values from reference + self.color_data.reference_dominant = ( + self.color_data.extract_dominant(ref_analysis_image, num_values=15) + ) + # Build its reference‐side map & blobs + self.color_data.create_map_with_blobs(ref_analysis_image, use_canvas=False) + + # Extract the 15 most dominant values from canvas + self.color_data.canvas_dominant = ( + self.color_data.extract_dominant(analysis_image, num_values=6) + ) + self.color_data.create_map_with_blobs(analysis_image, use_canvas=True) + + # Create colors maps and blob information + self.color_data.create_map_with_blobs(ref_analysis_image, use_canvas=False) + self.color_data.create_map_with_blobs(analysis_image, use_canvas=True) + + # Match canvas and reference colors + self.match_values(is_color_analysis=True) + + # Update UI + self.update_pairs( + self.color_data, + self.color_pairs_container_layout, + lambda rgb: color_conversion.rgb_to_hsv(rgb), # convert RGB→HSV for display + self.show_pair_regions_color + ) + + # Create and display the initial overview with all matched pairs + self.show_all_matched_pairs(True) + text_to_display = "✅ Found dominant colors. Click on any value pair to see detailed comparison." + self.color_feedback_label.setText(text_to_display) + self.append_log_entry("Get color feedback", "Requested color feedback analysis") + + def show_all_matched_pairs(self, is_color_analysis=False): + """Show all matched pairs together - reference with canvas regions and canvas with reference regions.""" + if is_color_analysis: + # Use color canvas and color reference + if self.color_filtered_canvas is None or self.color_reference_image is None: + return + + ref_with_regions = self.color_reference_image.copy() + canvas_with_regions = self.color_filtered_canvas.copy() + matched_pairs = self.color_data.matched_pairs + + # Store canvas for region display + canvas_for_mask = self.color_filtered_canvas + ref_for_mask = self.color_reference_image + else: + # Use grayscale canvas and grayscale reference + if self.filtered_canvas is None or self.filtered_image is None: + return + + ref_with_regions = image_conversion._to_bgr(self.filtered_image) + canvas_with_regions = image_conversion._to_bgr(self.filtered_canvas) + matched_pairs = self.value_data.matched_pairs + + # Store canvas for region display + canvas_for_mask = self.filtered_canvas + ref_for_mask = self.filtered_image + + # Store 5 colors for getting all matched pairs + colors = [(255, 0, 0), (0, 255, 0), (0, 0, 255), (255, 255, 0), (128, 0, 128)] + + # Draw all matched pairs + i = 0 + for canvas_hex, ref_hex in matched_pairs.items(): + # region masks + cur_color = colors[i % len(colors)] + canvas_mask = self.get_region_mask(canvas_for_mask, canvas_hex, is_color_analysis) + if (is_color_analysis): + ref_mask = self.get_region_mask(ref_for_mask, ref_hex, is_color_analysis) + else: + ref_mask = self.get_region_mask(self.filtered_image, ref_hex, is_color_analysis) + + # Contours for both masks + canvas_contours, _ = cv2.findContours(canvas_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + ref_contours, _ = cv2.findContours(ref_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + + if (is_color_analysis): + # Filter contours by minimum area + min_area = 100 # Adjust this threshold as needed + canvas_contours = [c for c in canvas_contours if cv2.contourArea(c) > min_area] + ref_contours = [c for c in ref_contours if cv2.contourArea(c) > min_area] + + for contour in canvas_contours: + if len(contour) > 2: + smoothed_contour = self.smooth_contour(contour) + cv2.polylines(canvas_with_regions, [smoothed_contour], + isClosed=True, color=cur_color, thickness=9) + + for contour in ref_contours: + if len(contour) > 2: + smoothed_contour = self.smooth_contour(contour) + cv2.polylines(ref_with_regions, [smoothed_contour], + isClosed=True, color=cur_color, thickness=9) + + i+=1 + + # Display the overlays - canvas on left, reference on right + self.display_split_view(canvas_with_regions, ref_with_regions, is_color_analysis) + + def display_split_view(self, left_image, right_image, is_color_analysis=False): + """Display two images side by side in the preview labels. + + Args: + left_image: First image to display (numpy array) - or canvas + right_image: Second image to display (numpy array) - or reference + is_color: Whether to handle as color images (True) or values (False) + """ + left_display = image_conversion._to_rgb_for_display(left_image) + right_display = image_conversion._to_rgb_for_display(right_image) + + left_label = self.color_left_preview_label if is_color_analysis else self.value_left_preview_label + right_label = self.color_right_preview_label if is_color_analysis else self.value_right_preview_label + + # Convert left and right image + left_height, left_width = left_display.shape[:2] + left_bytes_per_line = left_display.shape[2] * left_width + left_qimage = QImage(left_display.data, left_width, left_height, + left_bytes_per_line, QImage.Format_RGB888) + left_pixmap = QPixmap.fromImage(left_qimage) + + right_height, right_width = right_display.shape[:2] + right_bytes_per_line = right_display.shape[2] * right_width + right_qimage = QImage(right_display.data, right_width, right_height, + right_bytes_per_line, QImage.Format_RGB888) + right_pixmap = QPixmap.fromImage(right_qimage) + + # Update the preview labels + left_label.setPixmap(left_pixmap.scaled( + left_label.width(), + left_label.height(), + Qt.KeepAspectRatio)) + + right_label.setPixmap(right_pixmap.scaled( + right_label.width(), + right_label.height(), + Qt.KeepAspectRatio)) + + self.export_filtered_image_as_png(left_display, f"feedback_canvas_{'color' if is_color_analysis else 'value'}") + self.export_filtered_image_as_png(right_display, f"feedback_reference_{'color' if is_color_analysis else 'value'}") + + def smooth_contour(self, contour, num_points=100): + """Smooth a contour using spline interpolation.""" + # Extract points from the contour + points = contour.reshape(-1, 2) + if len(points) < 3: # Need at least 3 points for smoothing + return contour + + x = points[:, 0] + y = points[:, 1] + + # Add the first point to the end to close the loop + x = np.append(x, x[0]) + y = np.append(y, y[0]) + + # Create a cumulative distance array for interpolation + t = np.zeros(len(x)) + t[1:] = np.sqrt((x[1:] - x[:-1])**2 + (y[1:] - y[:-1])**2) + t = np.cumsum(t) + + if t[-1] == 0: + return contour + t /= t[-1] + + # Generate new points using spline interpolation + t_new = np.linspace(0, 1, num_points) + x_new = np.interp(t_new, t, x) + y_new = np.interp(t_new, t, y) + + # Combine into a list of points + smoothed_points = np.array([x_new, y_new]).T.astype(np.int32) + return smoothed_points.reshape(-1, 1, 2) # Reshape for polylines function + + def match_values(self, is_color_analysis=False): + """ + Match canvas values to reference values based on spatial and color information. + Only consider values that have actual blob regions on the screen. + + Args: + is_color_analysis: Whether to match colors (True) or values (False) + """ + if is_color_analysis: + self.color_data.matched_pairs = self.match_values_generic( + self.color_data.canvas_dominant, + self.color_data.reference_dominant, + self.color_data.canvas_blobs, + self.color_data.reference_blobs, + True # is_color_analysis + ) + else: + self.value_data.matched_pairs = self.match_values_generic( + self.value_data.canvas_dominant, + self.value_data.reference_dominant, + self.value_data.canvas_blobs, + self.value_data.reference_blobs, + False # is_color_analysis + ) + + def match_values_generic(self, canvas_values, reference_values, canvas_blobs, reference_blobs, is_color_analysis): + """ + Generic function to match canvas values to reference values based on spatial and color information. + Only consider values that have actual blob regions on the screen. + + Args: + canvas_values: List of (value/rgb, hex_code) tuples for canvas + reference_values: List of (value/rgb, hex_code) tuples for reference + canvas_blobs: Dictionary of blob information for canvas + reference_blobs: Dictionary of blob information for reference + is_color_analysis: Whether this is color analysis (True) or value analysis (False) + + Returns: + Dictionary of matched pairs {canvas_hex: reference_hex} + """ + matched_pairs = {} + + canvas_values_with_blobs = [] + for (value, hex_code) in canvas_values: + blob = canvas_blobs.get(hex_code) + if blob and blob.points: + canvas_values_with_blobs.append((value, hex_code)) + + reference_values_with_blobs = [] + for (value, hex_code) in reference_values: + blob = reference_blobs.get(hex_code) + if blob and blob.points: + reference_values_with_blobs.append((value, hex_code)) + + if not canvas_values_with_blobs or not reference_values_with_blobs: + return matched_pairs + + # Calculate similarity matrix between all canvas and reference values with blobs + similarity_matrix = [] + + for c_value, c_hex in canvas_values_with_blobs: + c_bbox = canvas_blobs[c_hex].bbox + row = [] + for r_value, r_hex in reference_values_with_blobs: + r_bbox = reference_blobs[r_hex].bbox + + # Calculate color similarity (weighted 50%) + color_similarity = matching_algo.calculate_color_similarity(c_hex, r_hex, is_color_analysis) * 0.50 + + # Calculate spatial similarity (weighted 50%) + spatial_similarity = matching_algo.calculate_bbox_overlap(c_bbox, r_bbox) * 0.50 + + # Total similarity is weighted sum + total_similarity = color_similarity + spatial_similarity + row.append((total_similarity, r_hex)) + + similarity_matrix.append((c_hex, row)) + + # Greedy matching algorithm for canvas to reference + for c_hex, similarities in similarity_matrix: + best_match = max(similarities, key=lambda x: x[0]) + best_similarity, best_ref_hex = best_match + + # Only match if similarity is above threshold + if best_similarity >= 0.01: + matched_pairs[c_hex] = best_ref_hex + + return matched_pairs + + def update_pairs(self, data, container_layout, transform, click_fn): + """ + Generic UI refresher for both value and color pairs. + + - data: either self.value_data or self.color_data + - container_layout: either self.value_pairs_container_layout or self.color_pairs_container_layout + - transform: a function(feature) → display_value (e.g. identity or rgb→hsv) + - click_fn: either self.show_pair_regions_value or self.show_pair_regions_color + """ + # Clear existing widgets + while container_layout.count(): + w = container_layout.takeAt(0).widget() + if w: + w.deleteLater() + + # Determine which selection attribute we use + is_color_tab = ("color" in container_layout.parent().objectName()) + selected_attr = "selected_color_pair" if is_color_tab else "selected_value_pair" + + # Reset selection + setattr(self, selected_attr, None) + + # Rebuild widgets + for canvas_hex, ref_hex in data.matched_pairs.items(): + + # Get feature values + left_feat = next(v for v, h in data.canvas_dominant if h == canvas_hex) + right_feat = next(v for v, h in data.reference_dominant if h == ref_hex) + + # Convert to display (identity or hsv) + left_disp = transform(left_feat) + right_disp = transform(right_feat) + + pair = ValuePairWidget(left_disp, canvas_hex, right_disp, ref_hex) + + # Clicking the color boxes triggers region display + pair.canvas_button.clicked.connect(lambda _, c=canvas_hex, r=ref_hex: click_fn(c, r)) + pair.ref_button.clicked.connect(lambda _, c=canvas_hex, r=ref_hex: click_fn(c, r)) + + # Click on widget → highlight only this one + def on_pair_clicked(p=pair, attr=selected_attr): + prev = getattr(self, attr) + if prev is not None: + prev.set_highlight(False) + + setattr(self, attr, p) + p.set_highlight(True) + + pair.clicked.connect(on_pair_clicked) + container_layout.addWidget(pair) + + def get_region_mask(self, image, hex_code, is_color_analysis=False): + """Get a binary mask for a specific value region.""" + if (is_color_analysis): + """Get a binary mask for a specific RGB value region.""" + # Convert hex code to RGB values + r = int(hex_code[1:3], 16) + g = int(hex_code[3:5], 16) + b = int(hex_code[5:7], 16) + + # Define threshold for RGB similarity + threshold = 20 + + # Create a temporary RGB copy for masking but don't modify original + # This is beacuse cv2 requires RGB for masking + rgb_image = image.copy() + if len(rgb_image.shape) == 2: # If grayscale, convert to RGB + rgb_image = cv2.cvtColor(rgb_image, cv2.COLOR_GRAY2RGB) + elif rgb_image.shape[2] == 4: # If RGBA, convert to RGB + rgb_image = cv2.cvtColor(rgb_image, cv2.COLOR_BGRA2RGB) + elif rgb_image.shape[2] == 3 and image is not self.color_canvas_image: # If BGR and not canvas, convert to RGB + rgb_image = cv2.cvtColor(rgb_image, cv2.COLOR_BGR2RGB) + + # Create lower and upper bounds for each channel + lower_bound = np.array([ + max(0, r - threshold), + max(0, g - threshold), + max(0, b - threshold) + ]) + upper_bound = np.array([ + min(255, r + threshold), + min(255, g + threshold), + min(255, b + threshold) + ]) + + # Create a binary mask where the pixel values are within the threshold range + mask = cv2.inRange(rgb_image, lower_bound, upper_bound) + else: + value = int(hex_code[1:3], 16) # Extract grayscale value from hex + threshold = 15 # Threshold for value similarity + + if len(image.shape) == 3: + if image.shape[2] == 4: # BGRA + gray_image = cv2.cvtColor(image, cv2.COLOR_BGRA2GRAY) + else: # BGR + gray_image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) + else: + gray_image = image + + mask = np.zeros_like(gray_image, dtype=np.uint8) + lower_bound = max(0, value - threshold) + upper_bound = min(255, value + threshold) + mask[(gray_image >= lower_bound) & (gray_image <= upper_bound)] = 255 + + return mask + + hue_ranges = [ + (0, 30, "red"), + (30, 90, "yellow"), + (90, 150, "green"), + (150, 210, "cyan"), + (210, 270, "blue"), + (270, 360, "magenta") + ] + + def get_significant_contours(self, mask: np.ndarray, min_area: int = 150): + """Return only those contours in `mask` whose area > min_area.""" + contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) + return [c for c in contours if cv2.contourArea(c) > min_area] + + def draw_contours(self, img: np.ndarray, contours: list, color: tuple, thickness: int = 9): + """Draw smoothed contours onto `img` in-place.""" + for cnt in contours: + if len(cnt) > 2: + sm = self.smooth_contour(cnt) + cv2.polylines(img, [sm], isClosed=True, color=color, thickness=thickness) + + def overlay_and_show(self, + canvas_img: np.ndarray, + ref_img: np.ndarray, + canvas_hex: str, + ref_hex: str, + is_color: bool): + """ + 1) mask → contours + 2) draw red on canvas, green on reference + 3) display split + 4) if color, save out overlays for Krita + """ + # Build masks + m_c = self.get_region_mask(canvas_img, canvas_hex, is_color) + m_r = self.get_region_mask(ref_img, ref_hex, is_color) + + # Pick only the big ones + ct_c = self.get_significant_contours(m_c) + ct_r = self.get_significant_contours(m_r) + + cv = (cv2.cvtColor(canvas_img, cv2.COLOR_GRAY2BGR) + if (not is_color and canvas_img.ndim == 2) + else canvas_img.copy()) + rv = (cv2.cvtColor(ref_img, cv2.COLOR_GRAY2BGR) + if (not is_color and ref_img.ndim == 2) + else ref_img.copy()) + + # Draw red / green + self.draw_contours(cv, ct_c, (0, 0, 255)) + self.draw_contours(rv, ct_r, (0, 255, 0)) + + self.display_split_view(cv, rv, is_color) + + # Save for Krita (only for color analysis) + if is_color: + temp_dir = os.path.join( + os.path.expanduser("~/Library/Application Support/krita/pykrita/artkrit") + if sys.platform=='darwin' + else os.path.expanduser("~/.local/share/krita/pykrita/artkrit"), + "temp" + ) + os.makedirs(temp_dir, exist_ok=True) + cv2.imwrite(os.path.join(temp_dir, "canvas_color_overlay.png"), cv) + cv2.imwrite(os.path.join(temp_dir, "reference_color_overlay.png"), rv) + + def show_pair_regions_color(self, canvas_hex, ref_hex): + """Highlight the selected color pair on canvas and reference, and show feedback.""" + if self.color_reference_image is None or self.color_canvas_image is None: + return + + # look up the two RGBs, bail if missing + c_rgb = next((rgb for rgb,h in self.color_data.canvas_dominant if h==canvas_hex), None) + r_rgb = next((rgb for rgb,h in self.color_data.reference_dominant if h==ref_hex), None) + if not (c_rgb and r_rgb): + return + + # feedback label + feedback = text_feedback.get_color_feedback(color_conversion.rgb_to_hsv(c_rgb), color_conversion.rgb_to_hsv(r_rgb), self.hue_ranges) + self.color_feedback_label.setText(feedback) + self.append_log_entry("color feedback", feedback) + + # overlay & show + self.overlay_and_show(self.color_canvas_image, self.color_reference_image, + canvas_hex, ref_hex, is_color=True) + + def show_pair_regions_value(self, canvas_hex, ref_hex): + """Highlight the selected value pair on canvas and reference, and show feedback.""" + if self.filtered_canvas is None or self.filtered_image is None: + return + + # feedback text (reuse existing get_value_feedback) + c_val = next((v for v,h in self.value_data.canvas_dominant if h==canvas_hex), None) + r_val = next((v for v,h in self.value_data.reference_dominant if h==ref_hex), None) + if c_val is not None and r_val is not None: + feedback = text_feedback.get_value_feedback(c_val, r_val) + self.value_feedback_label.setText(feedback) + self.append_log_entry("value feedback", feedback) + + # do exactly the same overlay→show, but in grayscale mode + self.overlay_and_show(self.filtered_canvas, self.filtered_image, + canvas_hex, ref_hex, is_color=False) + + def canvasChanged(self, event=None): + """This method is called whenever the canvas changes. It's required for all DockWidget subclasses in Krita.""" + pass + + def selectColor(self): + """Open the HS color picker and apply the chosen color to Krita's foreground.""" + dialog = CustomHSColorPickerDialog(self) + dialog.exec_() + + selectedColor = dialog.selectedColor() + if selectedColor.isValid(): + print(f"Selected Color: {selectedColor.name()}") + self.colorButton.setText(f"Color: {selectedColor.name()}") + + if Krita.instance().activeWindow() and Krita.instance().activeWindow().activeView(): + view = Krita.instance().activeWindow().activeView() + managedColor = ManagedColor("RGBA", "U8", "") + managedColor.setComponents([ + selectedColor.blueF(), + selectedColor.greenF(), + selectedColor.redF(), + 1.0 + ]) + view.setForeGroundColor(managedColor) + self.append_log_entry("color picker", f"Foreground color set to {selectedColor.name()}") + + def activateLassoTool(self): + """Activate Krita's lasso selection tool and prepare the fill options UI.""" + krita_instance = Krita.instance() + + action = krita_instance.action('KisToolSelectContiguous') + if action: + action.trigger() + self.lassoButton.setStyleSheet("background-color: #AED6F1;") + + self.fillGroup.setVisible(True) + + self.selectionTimer.start(500) + + QTimer.singleShot(500, lambda: self.lassoButton.setStyleSheet("")) + + + def selectFillColor(self): + """Open the HS picker seeded by the current selection's average value.""" + # Extract the dominant value from the selection + doc = Krita.instance().activeDocument() + if doc: + selection = doc.selection() + if selection: + node = doc.activeNode() + average_value = self.extractAverageValueFromSelection(node, selection) + if average_value is not None: + dialog = CustomHSColorPickerDialog(self, average_value) + dialog.exec_() + + selectedColor = dialog.selectedColor() + if selectedColor.isValid(): + self.currentFillColor = selectedColor + self.fillColorButton.setStyleSheet(f"background-color: {selectedColor.name()};") + + + def fillSelection(self): + """Fill the current selection with the selected fill color.""" + krita_instance = Krita.instance() + doc = krita_instance.activeDocument() + selection = doc.selection() + node = doc.activeNode() + average_value = self.extractAverageValueFromSelection(node, selection) + if average_value is None: + print("Failed to extract average value from selection") + return + + # Get the selected H and S from the color picker + selected_hue = self.currentFillColor.hue() + selected_saturation = self.currentFillColor.saturation() + + # Create the new color with the extracted value and selected H and S + new_color = QColor.fromHsv(selected_hue, selected_saturation, average_value) + print(f"New fill color: {new_color.name()}") + + # Convert QColor to ManagedColor for Krita + managedColor = ManagedColor("RGBA", "U8", "") + managedColor.setComponents([ + new_color.blueF(), + new_color.greenF(), + new_color.redF(), + 1.0 # Fully opaque + ]) + + # Set the foreground color + if krita_instance.activeWindow() and krita_instance.activeWindow().activeView(): + view = krita_instance.activeWindow().activeView() + view.setForeGroundColor(managedColor) + print("Foreground color set to new color") + + # Trigger the fill tool action + fillToolAction = krita_instance.action('KritaFill/KisToolFill') + if fillToolAction: + print("Triggering fill tool action") + fillToolAction.trigger() + + QTimer.singleShot(100, lambda: self.triggerFillForeground(krita_instance)) + + else: + print("Could not find fill tool action") + + def triggerFillForeground(self, krita_instance): + """ + Triggers the fill_foreground action after the fill tool is activated. + """ + fillAction = krita_instance.action('fill_foreground') + if fillAction: + fillAction.trigger() + + def extractAverageValueFromSelection(self, node, selection): + """ + Extracts the dominant brightness (value) from the selected area using the HSV color space. + """ + try: + print("Extracting pixel data from selection...") + + # Get the pixel data from the selected area + pixel_data = node.projectionPixelData( + selection.x(), selection.y(), selection.width(), selection.height() + ).data() + + pixels = [] + for i in range(0, len(pixel_data), 4): + r = pixel_data[i] + g = pixel_data[i + 1] + b = pixel_data[i + 2] + pixels.append((r, g, b)) + + # Calculate the frequency of each brightness (value) level + value_counts = {} + for r, g, b in pixels: + # Convert RGB to HSV + hsv_color = QColor(r, g, b).getHsv() + value = hsv_color[2] + + # Count the frequency of each value + if value in value_counts: + value_counts[value] += 1 + else: + value_counts[value] = 1 + + # Find the dominant value (the one with the highest frequency) + dominant_value = max(value_counts, key=value_counts.get) + print(f"Dominant value (brightness): {dominant_value}") + + return dominant_value + + except Exception as e: + print(f"Error extracting dominant value: {str(e)}") + return None + + def checkSelection(self): + doc = Krita.instance().activeDocument() + if doc: + selection = doc.selection() + if selection: + self.fillButton.setEnabled(True) + else: + self.fillButton.setEnabled(False) + + if self.fillGroup.isVisible(): + self.selectionTimer.start(500) + + def zoom_in(self): + """Zoom in on the image""" + self.image_label.setMinimumSize( + int(self.image_label.minimumWidth() * 1.2), + int(self.image_label.minimumHeight() * 1.2) + ) + self.image_label.setMaximumSize( + int(self.image_label.maximumWidth() * 1.2), + int(self.image_label.maximumHeight() * 1.2) + ) + self.image_label.update() + self.append_log_entry("zoom in", "Zoomed in on image preview") + + def zoom_out(self): + """Zoom out on the image""" + self.image_label.setMinimumSize( + int(self.image_label.minimumWidth() / 1.2), + int(self.image_label.minimumHeight() / 1.2) + ) + self.image_label.setMaximumSize( + int(self.image_label.maximumWidth() / 1.2), + int(self.image_label.maximumHeight() / 1.2) + ) + self.image_label.update() + self.append_log_entry("zoom out", "Zoomed out on image preview") + + def get_canvas_data(self): + """Get the current canvas data as a numpy array.""" + document = Krita.instance().activeDocument() + if not document: + return None + + active_layer = document.activeNode() + doc_width, doc_height = document.width(), document.height() + pixel_data = active_layer.pixelData(0, 0, doc_width, doc_height) + pixel_array = np.frombuffer(pixel_data, dtype=np.uint8).reshape(doc_height, doc_width, -1) + + # Downsample pixel array to half size using cv2.resize + pixel_array = cv2.resize(pixel_array, (doc_width//2, doc_height//2), interpolation=cv2.INTER_AREA) + + return pixel_array + + def show_current_canvas(self): + """Show the current canvas data in grayscale in the left preview.""" + pixel_array = self.get_canvas_data() + if pixel_array is None: + self.value_feedback_label.setText("⚠️ No document is open") + return + + self.value_canvas_image = image_conversion._to_grayscale(pixel_array) + + # Apply default Gaussian filter if no filter is selected + if not self.current_filter: + self.gaussian_radio.setChecked(True) + self.current_filter = "gaussian" + self.slider_label.show() + self.slider.show() + + # Apply the filter to create filtered_canvas + self.update_preview() + + # Display the grayscale image + self.display_preview(self.value_canvas_image, False) + self.value_feedback_label.setText("✅ Showing current canvas in grayscale") + self.append_log_entry("set canvas for value", "Set current canvas for value analysis") + + +class CustomHSColorPickerDialog(QDialog): + def __init__(self, parent=None, extracted_value=None): + super().__init__(parent) + + self.setWindowTitle("Select Color") + self.setModal(True) + + self.currentColor = QColor(255, 0, 0) + self.currentHue = 0 + self.currentSaturation = 255 + self.currentValue = extracted_value # Use the extracted value (if provided) + + self.huePicker = HuePicker(self) + self.saturationValuePicker = SaturationValuePicker(self, extracted_value) + + self.colorPreview = QLabel() + self.colorPreview.setFixedSize(100, 100) + self.colorPreview.setStyleSheet(f"background-color: {self.currentColor.name()};") + + self.okButton = QPushButton("OK") + self.okButton.clicked.connect(self.accept) + + self.cancelButton = QPushButton("Cancel") + self.cancelButton.clicked.connect(self.reject) + + layout = QVBoxLayout() + layout.addWidget(QLabel("Select Hue")) + layout.addWidget(self.huePicker) + layout.addWidget(QLabel("Select Saturation")) + layout.addWidget(self.saturationValuePicker) + layout.addWidget(self.colorPreview) + layout.addWidget(self.okButton) + layout.addWidget(self.cancelButton) + + self.setLayout(layout) + + self.huePicker.colorChanged.connect(self.updateFromHue) + self.saturationValuePicker.colorChanged.connect(self.updateColor) + + def updateFromHue(self): + self.currentHue = self.huePicker.getHue() + self.saturationValuePicker.setHue(self.currentHue) + self.updateColor() + + def updateColor(self): + self.currentColor = self.saturationValuePicker.getColor() + self.colorPreview.setStyleSheet(f"background-color: {self.currentColor.name()};") + + def selectedColor(self): + return self.currentColor + +class HuePicker(QWidget): + colorChanged = pyqtSignal() + + def __init__(self, parent=None): + super().__init__(parent) + self.setFixedSize(200, 200) + self.setStyleSheet("background-color: #f1f1f1; border: 1px solid #ccc;") + + self.hue = 0 + self.setAutoFillBackground(True) + + def paintEvent(self, event): + painter = QPainter(self) + rect = self.rect() + + outer_radius = min(rect.width(), rect.height()) / 2 + inner_radius = outer_radius - 20 + + gradient = QConicalGradient(rect.center(), 90) + for i in range(360): + gradient.setColorAt(i / 360, QColor.fromHsv(i, 255, 255)) + + path = QPainterPath() + path.addEllipse(rect.center(), outer_radius, outer_radius) + path.addEllipse(rect.center(), inner_radius, inner_radius) + + painter.setBrush(QBrush(gradient)) + painter.setPen(Qt.NoPen) + painter.drawPath(path) + + def mousePressEvent(self, event): + if event.button() == Qt.LeftButton: + self.updateHue(event.pos()) + + def mouseMoveEvent(self, event): + if event.buttons() == Qt.LeftButton: + self.updateHue(event.pos()) + + def updateHue(self, pos): + center = self.rect().center() + dx = pos.x() - center.x() + dy = pos.y() - center.y() + angle = math.degrees(math.atan2(dy, dx)) + self.hue = int((360 - ((angle + 90 + 360) % 360)) % 360) + + self.update() + self.colorChanged.emit() + + def getHue(self): + return self.hue + +class SaturationValuePicker(QWidget): + colorChanged = pyqtSignal() + + def __init__(self, parent=None, extracted_value=None): + super().__init__(parent) + self.setFixedSize(200, 50 if extracted_value is not None else 200) + self.setStyleSheet("background-color: #f1f1f1; border: 1px solid #ccc;") + + self.hue = 0 + self.saturation = 255 + self.value = extracted_value # Use the extracted value (if provided) + self.extracted_value = extracted_value + self.setAutoFillBackground(True) + + def setHue(self, hue): + """ + Set the hue for the color range. + """ + self.hue = hue + self.update() + + def paintEvent(self, event): + """ + Draw either a full saturation-value square or a horizontal slider for saturation. + """ + painter = QPainter(self) + rect = self.rect() + + if self.extracted_value is None: + # Draw the full saturation-value square + for x in range(rect.width()): + for y in range(rect.height()): + saturation = int((x / rect.width()) * 255) + value = int((y / rect.height()) * 255) + color = QColor.fromHsv(self.hue, saturation, value) + painter.setPen(color) + painter.drawPoint(x, y) + else: + # Draw a horizontal slider for saturation (value is fixed) + for x in range(rect.width()): + saturation = int((x / rect.width()) * 255) + color = QColor.fromHsv(self.hue, saturation, self.extracted_value) + painter.setPen(color) + painter.drawLine(x, 0, x, rect.height()) + + def mousePressEvent(self, event): + if event.button() == Qt.LeftButton: + self.updateColorFromPosition(event.pos()) + + def mouseMoveEvent(self, event): + if event.buttons() == Qt.LeftButton: + self.updateColorFromPosition(event.pos()) + + def updateColorFromPosition(self, pos): + """ + Update the selected color based on the mouse position. + """ + rect = self.rect() + x = pos.x() + + # Calculate the saturation based on the x position + self.saturation = int((x / rect.width()) * 255) + + # If no extracted value, calculate the value based on the y position + if self.extracted_value is None: + y = pos.y() + self.value = int((y / rect.height()) * 255) + + # Emit the color change signal + self.colorChanged.emit() + + def getColor(self): + """ + Get the currently selected color. + """ + return QColor.fromHsv(self.hue, self.saturation, self.value) \ No newline at end of file