Merge pull request #49 from outis1one/claude/awesome-bohr-mks3co
Fix startup crash: replace sklearn k-means with pure numpy implementa…
This commit is contained in:
@@ -511,8 +511,6 @@ async def enhance(req: EnhanceRequest):
|
|||||||
|
|
||||||
# ─── Extract colors ───────────────────────────────────────────────────────────
|
# ─── Extract colors ───────────────────────────────────────────────────────────
|
||||||
|
|
||||||
from sklearn.cluster import KMeans as _KMeans
|
|
||||||
|
|
||||||
class ExtractColorsRequest(BaseModel):
|
class ExtractColorsRequest(BaseModel):
|
||||||
image: str # base64
|
image: str # base64
|
||||||
count: int = 6
|
count: int = 6
|
||||||
@@ -520,28 +518,49 @@ class ExtractColorsRequest(BaseModel):
|
|||||||
|
|
||||||
def _extract_colors(image_bytes: bytes, count: int) -> list[str]:
|
def _extract_colors(image_bytes: bytes, count: int) -> list[str]:
|
||||||
"""
|
"""
|
||||||
Resize image to 150×150, k-means cluster pixels into `count` groups,
|
Resize image to 150×150, k-means cluster pixels into `count` groups
|
||||||
sort by cluster size (largest first), return as hex strings.
|
using pure numpy (no sklearn dependency), return hex strings by frequency.
|
||||||
"""
|
"""
|
||||||
|
import numpy as np
|
||||||
|
from PIL import Image
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
count = max(1, min(count, 32))
|
count = max(1, min(count, 32))
|
||||||
|
|
||||||
pil = _Image.open(_io.BytesIO(image_bytes)).convert("RGB").resize((150, 150))
|
pil = Image.open(BytesIO(image_bytes)).convert("RGB").resize((150, 150))
|
||||||
pixels = _np.array(pil, dtype=_np.float32).reshape(-1, 3) # (N, 3)
|
pixels = np.array(pil, dtype=np.float32).reshape(-1, 3) # (22500, 3)
|
||||||
|
n = len(pixels)
|
||||||
|
|
||||||
km = _KMeans(n_clusters=count, n_init=10, random_state=42)
|
# Initialise centers with k-means++ seeding
|
||||||
labels = km.fit_predict(pixels)
|
rng = np.random.default_rng(42)
|
||||||
centers = km.cluster_centers_ # (count, 3)
|
centers = [pixels[rng.integers(n)]]
|
||||||
|
for _ in range(count - 1):
|
||||||
|
dists = np.min([np.sum((pixels - c) ** 2, axis=1) for c in centers], axis=0)
|
||||||
|
probs = dists / dists.sum()
|
||||||
|
centers.append(pixels[rng.choice(n, p=probs)])
|
||||||
|
centers = np.array(centers)
|
||||||
|
|
||||||
# Count pixels per cluster and sort by frequency descending
|
labels = np.zeros(n, dtype=np.int32)
|
||||||
counts = _np.bincount(labels, minlength=count)
|
for _ in range(20): # max 20 iterations
|
||||||
order = _np.argsort(-counts) # descending
|
# Assign each pixel to nearest center
|
||||||
|
dists = np.sum((pixels[:, None] - centers[None]) ** 2, axis=2) # (n, k)
|
||||||
|
new_labels = np.argmin(dists, axis=1)
|
||||||
|
if np.all(new_labels == labels):
|
||||||
|
break
|
||||||
|
labels = new_labels
|
||||||
|
# Recompute centers
|
||||||
|
for k in range(count):
|
||||||
|
mask = labels == k
|
||||||
|
if mask.any():
|
||||||
|
centers[k] = pixels[mask].mean(axis=0)
|
||||||
|
|
||||||
hex_colors = []
|
counts = np.bincount(labels, minlength=count)
|
||||||
for idx in order:
|
order = np.argsort(-counts)
|
||||||
r, g, b = centers[idx].astype(int).clip(0, 255)
|
|
||||||
hex_colors.append(f"#{r:02x}{g:02x}{b:02x}")
|
|
||||||
|
|
||||||
return hex_colors
|
return [
|
||||||
|
"#{:02x}{:02x}{:02x}".format(*centers[i].astype(int).clip(0, 255))
|
||||||
|
for i in order
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
@router.post("/extract-colors")
|
@router.post("/extract-colors")
|
||||||
|
|||||||
Reference in New Issue
Block a user