Download tiling.py from enginil/ClearRealityV1-CoreAI: direct link, hf CLI and curl.
- Browser
- Download file 4.66 kB
-
https://huggingface.co/enginil/ClearRealityV1-CoreAI/resolve/main/tiling.py
- Command line
-
hf download hf://enginil/ClearRealityV1-CoreAI/tiling.py
-
curl -L -o tiling.py https://huggingface.co/enginil/ClearRealityV1-CoreAI/resolve/main/tiling.py
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") | |