mirror of
https://github.com/tiennm99/ArtKrit.git
synced 2026-10-03 05:18:45 +00:00
Merge remote-tracking branch 'origin/Asya' into catherine
This commit is contained in:
commit
627b3f73cf
16 files changed
+6029
No files matched your search
@@ -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
|
||||
@@ -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.
|
||||
|
||||
<img width="1062" height="618" alt="fig_teaser" src="https://github.com/user-attachments/assets/3bae40d9-9e91-4251-95b7-b7d251f5a3b5" />
|
||||
|
||||
**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)
|
||||
@@ -0,0 +1 @@
|
||||
from .artkrit import *
|
||||
+1078
File diff suppressed because it is too large.
Load diff
@@ -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
|
||||
File diff suppressed because it is too large.
Load diff
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
@@ -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
|
||||
File diff suppressed because it is too large.
Load diff
Reference in new issue
Block a user