Merge remote-tracking branch 'origin/Asya' into catherine

This commit is contained in:
cliu26 committed 2025-11-24 14:23:04 -08:00
commit 627b3f73cf
16 files changed
+6029

No files matched your search

+177
View File
@@ -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
+92
View File
@@ -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)
+1
View File
@@ -0,0 +1 @@
from .artkrit import *
+1078
View File
File diff suppressed because it is too large. Load diff
+11
View File
@@ -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
+545
View File
@@ -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
+187
View File
@@ -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)
+254
View File
@@ -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