#!/usr/bin/env python3.11 from __future__ import annotations from pathlib import Path from typing import Callable import numpy as np from PIL import Image def tile_positions(length: int, tile: int, overlap: int) -> list[int]: if tile <= 0: raise ValueError("tile must be > 0") if overlap < 0 or overlap >= tile: raise ValueError("overlap must satisfy 0 <= overlap < tile") if length <= tile: return [0] stride = tile - overlap last = length - tile pos = list(range(0, last + 1, stride)) if pos[-1] != last: pos.append(last) return pos def keep_interval( positions: list[int], i: int, tile: int, full_length: int ) -> tuple[int, int]: """Return the global interval owned by tile i. Adjacent tiles meet exactly at the midpoint of their actual overlap. This also handles the shorter/non-standard final overlap correctly. """ p = positions[i] if i == 0: start = 0 else: prev = positions[i - 1] # midpoint of overlap [p, prev + tile) start = (p + (prev + tile)) // 2 if i == len(positions) - 1: end = full_length else: nxt = positions[i + 1] # midpoint of overlap [nxt, p + tile) end = (nxt + (p + tile)) // 2 return start, end def pad_tile_edge( tile_arr: np.ndarray, target_h: int, target_w: int ) -> tuple[np.ndarray, int, int]: h, w = tile_arr.shape[:2] pad_h = max(0, target_h - h) pad_w = max(0, target_w - w) if pad_h or pad_w: tile_arr = np.pad( tile_arr, ((0, pad_h), (0, pad_w), (0, 0)), mode="edge", ) return tile_arr, h, w def tiled_upscale( image_rgb01: np.ndarray, infer_tile: Callable[[np.ndarray], np.ndarray], *, scale: int, tile: int = 512, overlap: int = 64, progress_prefix: str = "", ) -> np.ndarray: """Upscale HWC RGB [0,1] using fixed-size tiles and midpoint-discard stitching. infer_tile receives exactly tile x tile x 3 float32 data and returns (tile*scale) x (tile*scale) x 3. """ if image_rgb01.ndim != 3 or image_rgb01.shape[2] != 3: raise ValueError(f"expected HWC RGB image, got {image_rgb01.shape}") h, w = image_rgb01.shape[:2] ys = tile_positions(h, tile, overlap) xs = tile_positions(w, tile, overlap) out = np.empty((h * scale, w * scale, 3), dtype=np.float32) total = len(ys) * len(xs) n = 0 for yi, y in enumerate(ys): gy0, gy1 = keep_interval(ys, yi, tile, h) for xi, x in enumerate(xs): gx0, gx1 = keep_interval(xs, xi, tile, w) n += 1 if progress_prefix: print( f"\r{progress_prefix} tile {n}/{total} " f"@ x={x}, y={y}", end="", flush=True, ) crop = image_rgb01[y : min(y + tile, h), x : min(x + tile, w)] crop, original_h, original_w = pad_tile_edge(crop, tile, tile) pred = infer_tile(crop.astype(np.float32, copy=False)) expected = (tile * scale, tile * scale, 3) if pred.shape != expected: raise RuntimeError( f"inference returned {pred.shape}, expected {expected}" ) # Drop padded output before applying ownership interval. pred = pred[: original_h * scale, : original_w * scale] ly0 = (gy0 - y) * scale ly1 = (gy1 - y) * scale lx0 = (gx0 - x) * scale lx1 = (gx1 - x) * scale out[ gy0 * scale : gy1 * scale, gx0 * scale : gx1 * scale, ] = pred[ly0:ly1, lx0:lx1] if progress_prefix: print() return out def load_rgb01(path: str | Path) -> np.ndarray: im = Image.open(Path(path).expanduser().resolve()).convert("RGB") return np.asarray(im, dtype=np.float32) / 255.0 def save_rgb01( arr: np.ndarray, path: str | Path, *, jpeg_quality: int = 75, jpeg_subsampling: int = 2, ) -> None: path = Path(path).expanduser().resolve() path.parent.mkdir(parents=True, exist_ok=True) u8 = np.rint(np.clip(arr, 0.0, 1.0) * 255.0).astype(np.uint8) im = Image.fromarray(u8, mode="RGB") suffix = path.suffix.lower() if suffix in {".jpg", ".jpeg"}: im.save( path, quality=jpeg_quality, subsampling=jpeg_subsampling, ) elif suffix == ".png": im.save(path) else: raise ValueError("output extension must be .png, .jpg, or .jpeg")