Spaces:
Build error
Build error
| import random | |
| import cv2 | |
| import numpy as np | |
| from PIL import Image, ImageEnhance, ImageFilter, ImageOps | |
| from augmenator.spatial import apply_spatial_ops | |
| from augmenator.style_transfer import STYLE_MODELS, STYLE_TAGS, apply_style | |
| ALLOWED_TAGS = { | |
| "brighten", | |
| "darken", | |
| "warmer", | |
| "cooler", | |
| "rotate", | |
| "rotate_90_random", | |
| "rotate_left", | |
| "rotate_right", | |
| "rotate_180", | |
| "flip", | |
| "flip_vertical", | |
| "blur", | |
| "sharpen", | |
| "saturate", | |
| "desaturate", | |
| "crop_zoom", | |
| "contrast_up", | |
| "contrast_down", | |
| "gamma_up", | |
| "gamma_down", | |
| "autocontrast", | |
| "posterize", | |
| "sepia", | |
| "hue_shift", | |
| "tint_red", | |
| "tint_green", | |
| "tint_blue", | |
| "invert", | |
| "perspective", | |
| } | set(STYLE_TAGS) | |
| def _clamp_strength(strength: float) -> float: | |
| return max(0.5, min(1.5, strength)) | |
| def _warm_cool_pil(image: Image.Image, warmer: bool, amount: float) -> Image.Image: | |
| r, g, b = image.split() | |
| factor = 1.0 + (0.25 * amount if warmer else -0.15 * amount) | |
| r = r.point(lambda p: min(255, int(p * factor))) | |
| b = b.point(lambda p: max(0, int(p * (2.0 - factor)))) | |
| return Image.merge("RGB", (r, g, b)) | |
| def _apply_sepia(image: Image.Image, amount: float) -> Image.Image: | |
| arr = np.array(image).astype(np.float32) | |
| r, g, b = arr[..., 0], arr[..., 1], arr[..., 2] | |
| tr = 0.393 * r + 0.769 * g + 0.189 * b | |
| tg = 0.349 * r + 0.686 * g + 0.168 * b | |
| tb = 0.272 * r + 0.534 * g + 0.131 * b | |
| blend = amount | |
| out = np.stack( | |
| [ | |
| r * (1 - blend) + tr * blend, | |
| g * (1 - blend) + tg * blend, | |
| b * (1 - blend) + tb * blend, | |
| ], | |
| axis=-1, | |
| ) | |
| return Image.fromarray(np.clip(out, 0, 255).astype(np.uint8)) | |
| def _apply_hue_shift(image: Image.Image, amount: float) -> Image.Image: | |
| arr = np.array(image.convert("RGB")) | |
| hsv = cv2.cvtColor(arr, cv2.COLOR_RGB2HSV).astype(np.float32) | |
| shift = int(18 * amount) | |
| hsv[..., 0] = (hsv[..., 0] + shift) % 180 | |
| rgb = cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2RGB) | |
| return Image.fromarray(rgb) | |
| def _apply_tint(image: Image.Image, channel: str, amount: float) -> Image.Image: | |
| arr = np.array(image).astype(np.float32) | |
| idx = {"red": 0, "green": 1, "blue": 2}[channel] | |
| boost = 1.0 + 0.35 * amount | |
| arr[..., idx] = np.clip(arr[..., idx] * boost, 0, 255) | |
| return Image.fromarray(arr.astype(np.uint8)) | |
| def _rotate_cardinal(image: Image.Image, degrees: int) -> Image.Image: | |
| mapping = { | |
| 90: Image.ROTATE_90, | |
| 180: Image.ROTATE_180, | |
| 270: Image.ROTATE_270, | |
| } | |
| key = degrees % 360 | |
| if key not in mapping: | |
| raise ValueError(f"Unsupported cardinal rotation: {degrees}") | |
| return image.transpose(mapping[key]) | |
| def _apply_gamma(image: Image.Image, gamma: float) -> Image.Image: | |
| arr = np.array(image.convert("RGB")).astype(np.float32) / 255.0 | |
| corrected = np.power(arr, gamma) | |
| return Image.fromarray(np.clip(corrected * 255.0, 0, 255).astype(np.uint8)) | |
| def _apply_tag(image: Image.Image, tag: str, strength: float) -> Image.Image: | |
| s = _clamp_strength(strength) | |
| if tag == "brighten": | |
| return ImageEnhance.Brightness(image).enhance(1.0 + 0.3 * s) | |
| if tag == "darken": | |
| return ImageEnhance.Brightness(image).enhance(1.0 - 0.25 * s) | |
| if tag == "warmer": | |
| return _warm_cool_pil(image, warmer=True, amount=s) | |
| if tag == "cooler": | |
| return _warm_cool_pil(image, warmer=False, amount=s) | |
| if tag == "contrast_up": | |
| return ImageEnhance.Contrast(image).enhance(1.0 + 0.4 * s) | |
| if tag == "contrast_down": | |
| return ImageEnhance.Contrast(image).enhance(1.0 - 0.3 * s) | |
| if tag == "gamma_up": | |
| gamma = 1.0 / (1.0 + 0.3 * s) | |
| return _apply_gamma(image, gamma) | |
| if tag == "gamma_down": | |
| gamma = 1.0 + 0.3 * s | |
| return _apply_gamma(image, gamma) | |
| if tag == "autocontrast": | |
| return ImageOps.autocontrast(image.convert("RGB"), cutoff=int(2 * s)) | |
| if tag == "posterize": | |
| bits = max(3, int(8 - s)) | |
| return ImageOps.posterize(image.convert("RGB"), bits) | |
| if tag == "sepia": | |
| return _apply_sepia(image, s) | |
| if tag == "hue_shift": | |
| return _apply_hue_shift(image, s) | |
| if tag == "tint_red": | |
| return _apply_tint(image, "red", s) | |
| if tag == "tint_green": | |
| return _apply_tint(image, "green", s) | |
| if tag == "tint_blue": | |
| return _apply_tint(image, "blue", s) | |
| if tag == "invert": | |
| return ImageOps.invert(image.convert("RGB")) | |
| if tag == "rotate": | |
| angle = random.uniform(0, 360) | |
| fill = (128, 128, 128) if image.mode == "RGB" else (128, 128, 128, 255) | |
| rotated = image.rotate(angle, expand=True, fillcolor=fill) | |
| return rotated, f"rotate({angle:.1f}°)" | |
| if tag == "rotate_90_random": | |
| degrees = random.choice([90, 270]) | |
| direction = "CCW" if degrees == 90 else "CW" | |
| return _rotate_cardinal(image, degrees), f"rotate_90({direction})" | |
| if tag == "rotate_left": | |
| return _rotate_cardinal(image, 90), "rotate_left(90° CCW)" | |
| if tag == "rotate_right": | |
| return _rotate_cardinal(image, 270), "rotate_right(90° CW)" | |
| if tag == "rotate_180": | |
| return _rotate_cardinal(image, 180), "rotate_180" | |
| if tag == "flip": | |
| return image.transpose(Image.FLIP_LEFT_RIGHT) | |
| if tag == "flip_vertical": | |
| return image.transpose(Image.FLIP_TOP_BOTTOM) | |
| if tag == "blur": | |
| return image.filter(ImageFilter.GaussianBlur(radius=1.5 * s)) | |
| if tag == "sharpen": | |
| return image.filter(ImageFilter.UnsharpMask(radius=2, percent=int(120 * s))) | |
| if tag == "saturate": | |
| return ImageEnhance.Color(image).enhance(1.0 + 0.4 * s) | |
| if tag == "desaturate": | |
| return ImageEnhance.Color(image).enhance(1.0 - 0.35 * s) | |
| if tag == "crop_zoom": | |
| w, h = image.size | |
| crop_ratio = max(0.6, 0.85 - 0.1 * (s - 1.0)) | |
| tw, th = int(w * crop_ratio), int(h * crop_ratio) | |
| left = (w - tw) // 2 | |
| top = (h - th) // 2 | |
| cropped = image.crop((left, top, left + tw, top + th)) | |
| return cropped.resize((w, h), Image.Resampling.LANCZOS) | |
| if tag == "perspective": | |
| return _apply_perspective(image, s) | |
| if tag in STYLE_TAGS: | |
| styled = apply_style(image, tag, strength) | |
| return styled, f"style({STYLE_MODELS[tag]['label']})" | |
| return image | |
| def _apply_perspective(image: Image.Image, strength: float) -> tuple[Image.Image, str]: | |
| """Perspective warp with output canvas fitted to include the full warped image.""" | |
| arr = np.array(image.convert("RGB")) | |
| height, width = arr.shape[:2] | |
| margin = 0.10 * strength * min(width, height) | |
| src = np.float32([[0, 0], [width - 1, 0], [width - 1, height - 1], [0, height - 1]]) | |
| offsets = np.random.uniform(-margin, margin, size=(4, 2)).astype(np.float32) | |
| dst = src + offsets | |
| matrix = cv2.getPerspectiveTransform(src, dst) | |
| warped_corners = cv2.perspectiveTransform(src.reshape(1, 4, 2), matrix).reshape(4, 2) | |
| x_min, y_min = warped_corners.min(axis=0) | |
| x_max, y_max = warped_corners.max(axis=0) | |
| translate = np.array([[1, 0, -x_min], [0, 1, -y_min], [0, 0, 1]], dtype=np.float32) | |
| full_matrix = translate @ matrix | |
| out_w = max(1, int(np.ceil(x_max - x_min))) | |
| out_h = max(1, int(np.ceil(y_max - y_min))) | |
| warped = cv2.warpPerspective( | |
| arr, | |
| full_matrix, | |
| (out_w, out_h), | |
| flags=cv2.INTER_LINEAR, | |
| borderMode=cv2.BORDER_REPLICATE, | |
| ) | |
| return Image.fromarray(warped), f"perspective(fitted {out_w}x{out_h})" | |
| def apply_augmentations( | |
| image: Image.Image, | |
| tags: list[str], | |
| strength: float = 1.0, | |
| spatial_ops: list[dict] | None = None, | |
| instruction: str = "", | |
| ) -> tuple[Image.Image, list[str], list[dict]]: | |
| if image.mode not in ("RGB", "RGBA"): | |
| image = image.convert("RGB") | |
| applied = [] | |
| applied_spatial = [] | |
| result = image.copy() | |
| # Spatial ops (cover/cutout) run first so OCR text boxes stay aligned | |
| # with the original image before geometric transforms like rotate. | |
| if spatial_ops: | |
| result, applied_spatial = apply_spatial_ops( | |
| result, spatial_ops, strength=strength, instruction=instruction | |
| ) | |
| for tag in tags: | |
| if tag not in ALLOWED_TAGS or tag in STYLE_TAGS: | |
| continue | |
| out = _apply_tag(result, tag, strength) | |
| if isinstance(out, tuple): | |
| result, label = out | |
| applied.append(label) | |
| else: | |
| result = out | |
| applied.append(tag) | |
| for tag in tags: | |
| if tag not in STYLE_TAGS: | |
| continue | |
| out = _apply_tag(result, tag, strength) | |
| if isinstance(out, tuple): | |
| result, label = out | |
| applied.append(label) | |
| else: | |
| result = out | |
| applied.append(tag) | |
| if result.mode == "RGBA": | |
| pass | |
| else: | |
| result = result.convert("RGB") | |
| return result, applied, applied_spatial | |