pi-codex-image-gen 0.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,440 @@
1
+ #!/usr/bin/env python3
2
+ """Remove a solid chroma-key background from an image.
3
+
4
+ This helper supports the imagegen skill's built-in-first transparent workflow:
5
+ generate an image on a flat key color, then convert that key color to alpha.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import argparse
11
+ from io import BytesIO
12
+ from pathlib import Path
13
+ import re
14
+ from statistics import median
15
+ import sys
16
+ from typing import Tuple
17
+
18
+
19
+ Color = Tuple[int, int, int]
20
+ KEY_DOMINANCE_THRESHOLD = 16.0
21
+ ALPHA_NOISE_FLOOR = 8
22
+
23
+
24
+ def _die(message: str, code: int = 1) -> None:
25
+ print(f"Error: {message}", file=sys.stderr)
26
+ raise SystemExit(code)
27
+
28
+
29
+ def _dependency_hint(package: str) -> str:
30
+ return (
31
+ "Activate the repo-selected environment first, then install it with "
32
+ f"`uv pip install {package}`. If this repo uses a local virtualenv, start with "
33
+ "`source .venv/bin/activate`; otherwise use this repo's configured shared fallback "
34
+ "environment."
35
+ )
36
+
37
+
38
+ def _load_pillow():
39
+ try:
40
+ from PIL import Image, ImageFilter
41
+ except ImportError:
42
+ _die(f"Pillow is required for chroma-key removal. {_dependency_hint('pillow')}")
43
+ return Image, ImageFilter
44
+
45
+
46
+ def _parse_key_color(raw: str) -> Color:
47
+ value = raw.strip()
48
+ match = re.fullmatch(r"#?([0-9a-fA-F]{6})", value)
49
+ if not match:
50
+ _die("key color must be a hex RGB value like #00ff00.")
51
+ hex_value = match.group(1)
52
+ return (
53
+ int(hex_value[0:2], 16),
54
+ int(hex_value[2:4], 16),
55
+ int(hex_value[4:6], 16),
56
+ )
57
+
58
+
59
+ def _validate_args(args: argparse.Namespace) -> None:
60
+ if args.tolerance < 0 or args.tolerance > 255:
61
+ _die("--tolerance must be between 0 and 255.")
62
+ if args.transparent_threshold < 0 or args.transparent_threshold > 255:
63
+ _die("--transparent-threshold must be between 0 and 255.")
64
+ if args.opaque_threshold < 0 or args.opaque_threshold > 255:
65
+ _die("--opaque-threshold must be between 0 and 255.")
66
+ if args.soft_matte and args.transparent_threshold >= args.opaque_threshold:
67
+ _die("--transparent-threshold must be lower than --opaque-threshold.")
68
+ if args.edge_feather < 0 or args.edge_feather > 64:
69
+ _die("--edge-feather must be between 0 and 64.")
70
+ if args.edge_contract < 0 or args.edge_contract > 16:
71
+ _die("--edge-contract must be between 0 and 16.")
72
+
73
+ src = Path(args.input)
74
+ if not src.exists():
75
+ _die(f"Input image not found: {src}")
76
+
77
+ out = Path(args.out)
78
+ if out.exists() and not args.force:
79
+ _die(f"Output already exists: {out} (use --force to overwrite)")
80
+
81
+ if out.suffix.lower() not in {".png", ".webp"}:
82
+ _die("--out must end in .png or .webp so the alpha channel is preserved.")
83
+
84
+
85
+ def _channel_distance(a: Color, b: Color) -> int:
86
+ return max(abs(a[0] - b[0]), abs(a[1] - b[1]), abs(a[2] - b[2]))
87
+
88
+
89
+ def _clamp_channel(value: float) -> int:
90
+ return max(0, min(255, int(round(value))))
91
+
92
+
93
+ def _smoothstep(value: float) -> float:
94
+ value = max(0.0, min(1.0, value))
95
+ return value * value * (3.0 - 2.0 * value)
96
+
97
+
98
+ def _soft_alpha(distance: int, transparent_threshold: float, opaque_threshold: float) -> int:
99
+ if distance <= transparent_threshold:
100
+ return 0
101
+ if distance >= opaque_threshold:
102
+ return 255
103
+ ratio = (float(distance) - transparent_threshold) / (
104
+ opaque_threshold - transparent_threshold
105
+ )
106
+ return _clamp_channel(255.0 * _smoothstep(ratio))
107
+
108
+
109
+ def _dominance_alpha(rgb: Color, key: Color) -> int:
110
+ spill_channels = _spill_channels(key)
111
+ if not spill_channels:
112
+ return 255
113
+
114
+ channels = [float(value) for value in rgb]
115
+ non_spill = [idx for idx in range(3) if idx not in spill_channels]
116
+ key_strength = (
117
+ min(channels[idx] for idx in spill_channels)
118
+ if len(spill_channels) > 1
119
+ else channels[spill_channels[0]]
120
+ )
121
+ non_key_strength = max((channels[idx] for idx in non_spill), default=0.0)
122
+ dominance = key_strength - non_key_strength
123
+ if dominance <= 0:
124
+ return 255
125
+
126
+ denominator = max(1.0, float(max(key)) - non_key_strength)
127
+ alpha = 1.0 - min(1.0, dominance / denominator)
128
+ return _clamp_channel(alpha * 255.0)
129
+
130
+
131
+ def _spill_channels(key: Color) -> list[int]:
132
+ key_max = max(key)
133
+ if key_max < 128:
134
+ return []
135
+ return [idx for idx, value in enumerate(key) if value >= key_max - 16 and value >= 128]
136
+
137
+
138
+ def _key_channel_dominance(rgb: Color, key: Color) -> float:
139
+ spill_channels = _spill_channels(key)
140
+ if not spill_channels:
141
+ return 0.0
142
+
143
+ channels = [float(value) for value in rgb]
144
+ non_spill = [idx for idx in range(3) if idx not in spill_channels]
145
+ key_strength = (
146
+ min(channels[idx] for idx in spill_channels)
147
+ if len(spill_channels) > 1
148
+ else channels[spill_channels[0]]
149
+ )
150
+ non_key_strength = max((channels[idx] for idx in non_spill), default=0.0)
151
+ return key_strength - non_key_strength
152
+
153
+
154
+ def _looks_key_colored(rgb: Color, key: Color, distance: int) -> bool:
155
+ if distance <= 32:
156
+ return True
157
+
158
+ spill_channels = _spill_channels(key)
159
+ if not spill_channels:
160
+ return True
161
+
162
+ return _key_channel_dominance(rgb, key) >= KEY_DOMINANCE_THRESHOLD
163
+
164
+
165
+ def _cleanup_spill(rgb: Color, key: Color, alpha: int = 255) -> Color:
166
+ if alpha >= 252:
167
+ return rgb
168
+
169
+ spill_channels = _spill_channels(key)
170
+ if not spill_channels:
171
+ return rgb
172
+
173
+ channels = [float(value) for value in rgb]
174
+ non_spill = [idx for idx in range(3) if idx not in spill_channels]
175
+ if non_spill:
176
+ anchor = max(channels[idx] for idx in non_spill)
177
+ cap = max(0.0, anchor - 1.0)
178
+ for idx in spill_channels:
179
+ if channels[idx] > cap:
180
+ channels[idx] = cap
181
+
182
+ return (
183
+ _clamp_channel(channels[0]),
184
+ _clamp_channel(channels[1]),
185
+ _clamp_channel(channels[2]),
186
+ )
187
+
188
+
189
+ def _apply_alpha_to_image(
190
+ image,
191
+ *,
192
+ key: Color,
193
+ tolerance: int,
194
+ spill_cleanup: bool,
195
+ soft_matte: bool,
196
+ transparent_threshold: float,
197
+ opaque_threshold: float,
198
+ ) -> int:
199
+ pixels = image.load()
200
+ width, height = image.size
201
+ transparent = 0
202
+
203
+ for y in range(height):
204
+ for x in range(width):
205
+ red, green, blue, alpha = pixels[x, y]
206
+ rgb = (red, green, blue)
207
+ distance = _channel_distance(rgb, key)
208
+ key_like = _looks_key_colored(rgb, key, distance)
209
+ output_alpha = (
210
+ min(
211
+ _soft_alpha(distance, transparent_threshold, opaque_threshold),
212
+ _dominance_alpha(rgb, key),
213
+ )
214
+ if soft_matte and key_like
215
+ else (0 if distance <= tolerance else 255)
216
+ )
217
+ output_alpha = int(round(output_alpha * (alpha / 255.0)))
218
+ if 0 < output_alpha <= ALPHA_NOISE_FLOOR:
219
+ output_alpha = 0
220
+
221
+ if output_alpha == 0:
222
+ pixels[x, y] = (0, 0, 0, 0)
223
+ transparent += 1
224
+ continue
225
+
226
+ if spill_cleanup and key_like:
227
+ red, green, blue = _cleanup_spill(rgb, key, output_alpha)
228
+ pixels[x, y] = (red, green, blue, output_alpha)
229
+
230
+ return transparent
231
+
232
+
233
+ def _contract_alpha(image, pixels: int):
234
+ if pixels == 0:
235
+ return image
236
+
237
+ _, ImageFilter = _load_pillow()
238
+ alpha = image.getchannel("A")
239
+ for _ in range(pixels):
240
+ alpha = alpha.filter(ImageFilter.MinFilter(3))
241
+ image.putalpha(alpha)
242
+ return image
243
+
244
+
245
+ def _apply_edge_feather(image, radius: float):
246
+ if radius == 0:
247
+ return image
248
+
249
+ _, ImageFilter = _load_pillow()
250
+ alpha = image.getchannel("A")
251
+ alpha = alpha.filter(ImageFilter.GaussianBlur(radius=radius))
252
+ image.putalpha(alpha)
253
+ return image
254
+
255
+
256
+ def _encode_image(image, output_format: str) -> bytes:
257
+ out = BytesIO()
258
+ image.save(out, format=output_format.upper())
259
+ return out.getvalue()
260
+
261
+
262
+ def _alpha_counts(image) -> tuple[int, int, int]:
263
+ pixels = image.load()
264
+ width, height = image.size
265
+ total = 0
266
+ transparent = 0
267
+ partial = 0
268
+
269
+ for y in range(height):
270
+ for x in range(width):
271
+ alpha = pixels[x, y][3]
272
+ total += 1
273
+ if alpha == 0:
274
+ transparent += 1
275
+ elif alpha < 255:
276
+ partial += 1
277
+
278
+ return total, transparent, partial
279
+
280
+
281
+ def _sample_border_key(image, mode: str) -> Color:
282
+ width, height = image.size
283
+ pixels = image.load()
284
+ samples: list[Color] = []
285
+
286
+ if mode == "corners":
287
+ patch = max(1, min(width, height, 12))
288
+ boxes = [
289
+ (0, 0, patch, patch),
290
+ (width - patch, 0, width, patch),
291
+ (0, height - patch, patch, height),
292
+ (width - patch, height - patch, width, height),
293
+ ]
294
+ for left, top, right, bottom in boxes:
295
+ for y in range(top, bottom):
296
+ for x in range(left, right):
297
+ red, green, blue = pixels[x, y][:3]
298
+ samples.append((red, green, blue))
299
+ else:
300
+ band = max(1, min(width, height, 6))
301
+ step = max(1, min(width, height) // 256)
302
+ for x in range(0, width, step):
303
+ for y in range(band):
304
+ red, green, blue = pixels[x, y][:3]
305
+ samples.append((red, green, blue))
306
+ red, green, blue = pixels[x, height - 1 - y][:3]
307
+ samples.append((red, green, blue))
308
+ for y in range(0, height, step):
309
+ for x in range(band):
310
+ red, green, blue = pixels[x, y][:3]
311
+ samples.append((red, green, blue))
312
+ red, green, blue = pixels[width - 1 - x, y][:3]
313
+ samples.append((red, green, blue))
314
+
315
+ if not samples:
316
+ _die("Could not sample background key color from image border.")
317
+
318
+ return (
319
+ int(round(median(sample[0] for sample in samples))),
320
+ int(round(median(sample[1] for sample in samples))),
321
+ int(round(median(sample[2] for sample in samples))),
322
+ )
323
+
324
+
325
+ def _remove_chroma_key(args: argparse.Namespace) -> None:
326
+ Image, _ = _load_pillow()
327
+ src = Path(args.input)
328
+ out = Path(args.out)
329
+
330
+ with Image.open(src) as image:
331
+ rgba = image.convert("RGBA")
332
+ key = (
333
+ _sample_border_key(rgba, args.auto_key)
334
+ if args.auto_key != "none"
335
+ else _parse_key_color(args.key_color)
336
+ )
337
+
338
+ transparent = _apply_alpha_to_image(
339
+ rgba,
340
+ key=key,
341
+ tolerance=args.tolerance,
342
+ spill_cleanup=args.spill_cleanup,
343
+ soft_matte=args.soft_matte,
344
+ transparent_threshold=args.transparent_threshold,
345
+ opaque_threshold=args.opaque_threshold,
346
+ )
347
+ rgba = _contract_alpha(rgba, args.edge_contract)
348
+ rgba = _apply_edge_feather(rgba, args.edge_feather)
349
+
350
+ total, transparent_after, partial_after = _alpha_counts(rgba)
351
+
352
+ out.parent.mkdir(parents=True, exist_ok=True)
353
+ output_format = "PNG" if out.suffix.lower() == ".png" else "WEBP"
354
+ out.write_bytes(_encode_image(rgba, output_format))
355
+
356
+ print(f"Wrote {out}")
357
+ print(f"Key color: #{key[0]:02x}{key[1]:02x}{key[2]:02x}")
358
+ print(f"Transparent pixels: {transparent_after}/{total}")
359
+ print(f"Partially transparent pixels: {partial_after}/{total}")
360
+ if transparent == 0:
361
+ print("Warning: no pixels matched the key color before feathering.", file=sys.stderr)
362
+
363
+
364
+ def _build_parser() -> argparse.ArgumentParser:
365
+ parser = argparse.ArgumentParser(
366
+ description="Remove a solid chroma-key background and write an image with alpha."
367
+ )
368
+ parser.add_argument("--input", required=True, help="Input image path.")
369
+ parser.add_argument("--out", required=True, help="Output .png or .webp path.")
370
+ parser.add_argument(
371
+ "--key-color",
372
+ default="#00ff00",
373
+ help="Hex RGB key color to remove, for example #00ff00.",
374
+ )
375
+ parser.add_argument(
376
+ "--tolerance",
377
+ type=int,
378
+ default=12,
379
+ help="Hard-key per-channel tolerance for matching the key color, 0-255.",
380
+ )
381
+ parser.add_argument(
382
+ "--auto-key",
383
+ choices=["none", "corners", "border"],
384
+ default="none",
385
+ help="Sample the key color from image corners or border instead of --key-color.",
386
+ )
387
+ parser.add_argument(
388
+ "--soft-matte",
389
+ action="store_true",
390
+ help="Use a smooth alpha ramp between transparent and opaque thresholds.",
391
+ )
392
+ parser.add_argument(
393
+ "--transparent-threshold",
394
+ type=float,
395
+ default=12.0,
396
+ help="Soft-matte distance at or below which pixels become fully transparent.",
397
+ )
398
+ parser.add_argument(
399
+ "--opaque-threshold",
400
+ type=float,
401
+ default=96.0,
402
+ help="Soft-matte distance at or above which pixels become fully opaque.",
403
+ )
404
+ parser.add_argument(
405
+ "--edge-feather",
406
+ type=float,
407
+ default=0.0,
408
+ help="Optional alpha blur radius for softened edges, 0-64.",
409
+ )
410
+ parser.add_argument(
411
+ "--edge-contract",
412
+ type=int,
413
+ default=0,
414
+ help="Shrink the visible alpha matte by this many pixels before feathering.",
415
+ )
416
+ parser.add_argument(
417
+ "--spill-cleanup",
418
+ dest="spill_cleanup",
419
+ action="store_true",
420
+ help="Reduce obvious key-color spill on opaque pixels.",
421
+ )
422
+ parser.add_argument(
423
+ "--despill",
424
+ dest="spill_cleanup",
425
+ action="store_true",
426
+ help="Alias for --spill-cleanup; decontaminate key-color edge spill.",
427
+ )
428
+ parser.add_argument("--force", action="store_true", help="Overwrite an existing output file.")
429
+ return parser
430
+
431
+
432
+ def main() -> None:
433
+ parser = _build_parser()
434
+ args = parser.parse_args()
435
+ _validate_args(args)
436
+ _remove_chroma_key(args)
437
+
438
+
439
+ if __name__ == "__main__":
440
+ main()