enginil's picture
Upload 9 files
15557f0 verified
Raw History Blame Contribute Delete
4.66 kB
#!/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")