"""Standalone inference utilities (numpy/PIL only — no torch needed).""" import numpy as np from PIL import Image, ImageDraw ALBEDO = slice(0, 3) DEPTH = slice(3, 4) NORMAL = slice(4, 7) SHADING = slice(7, 8) def resize_pad(img_rgb: Image.Image, image_size: int) -> Image.Image: """Scale + center-crop to a square (RGB in, RGB out).""" img_rgb = img_rgb.convert("RGB") w, h = img_rgb.size size = int(image_size) scale = max(size / float(w), size / float(h)) img_r = img_rgb.resize((max(1, int(round(w * scale))), max(1, int(round(h * scale)))), Image.BICUBIC) left = (img_r.size[0] - size) // 2 top = (img_r.size[1] - size) // 2 return img_r.crop((left, top, left + size, top + size)) def pil_to_np(img_rgb: Image.Image) -> np.ndarray: """PIL RGB -> [1, 3, H, W] float32 in [-1, 1].""" arr = np.array(img_rgb.convert("RGB"), dtype=np.float32).transpose(2, 0, 1) return arr[np.newaxis] / 255.0 * 2.0 - 1.0 def _to_01(output: np.ndarray) -> np.ndarray: return np.clip((output + 1.0) / 2.0, 0, 1) def _tile(arr01: np.ndarray, ts: int) -> Image.Image: return Image.fromarray( (arr01 * 255.0 + 0.5).astype("uint8")).resize((ts, ts), Image.BICUBIC) def _resize_map(out: np.ndarray, size: int) -> np.ndarray: """Bilinear-resize a [1, C, H, W] float map to [1, C, size, size].""" c = out.shape[1] res = np.empty((1, c, size, size), dtype=np.float32) for i in range(c): im = Image.fromarray(out[0, i].astype(np.float32), mode="F") res[0, i] = np.asarray(im.resize((size, size), Image.BILINEAR)) return res ENS_SCALES = (0.875, 1.0, 1.125) def run_ensemble_onnx(sess, img_rgb: Image.Image, image_size: int = 512, scales=ENS_SCALES) -> np.ndarray: """3-pass multi-scale median (torch-free). Returns [1, 8, S, S].""" outs = [] for s in scales: side = max(16, int(round(image_size * s)) // 16 * 16) # dict needs /16 scaled = resize_pad(img_rgb, side) x = pil_to_np(scaled).astype(np.float32) o = sess.run(None, {"input_rgb": x})[0] outs.append(o if side == image_size else _resize_map(o, image_size)) return np.median(np.stack(outs, axis=0), axis=0).astype(np.float32) def extract_maps(output: np.ndarray, img_rgb: Image.Image, ts: int = 256) -> dict: """Split a [1, 8, H, W] output ([-1, 1]) into labelled PIL tiles.""" assert output.shape[1] == 8, f"want 8 channels, got {output.shape[1]}" o = _to_01(output) alb = o[:, ALBEDO][0].transpose(1, 2, 0) sh = o[:, SHADING][0, 0] d = o[:, DEPTH][0, 0] span = d.max() - d.min() dg = (d - d.min()) / (span if span > 1e-6 else 1.0) return { "input": img_rgb.resize((ts, ts), Image.BICUBIC), "albedo": _tile(alb, ts), "shading": _tile(np.repeat(sh[..., None], 3, axis=2), ts), "depth": _tile(dg, ts).convert("RGB"), "normal": _tile(o[:, NORMAL][0].transpose(1, 2, 0), ts), "recon": _tile(np.clip(alb * sh[..., None], 0, 1), ts), } def build_grid(img_rgb: Image.Image, output: np.ndarray, ts: int = 256) -> Image.Image: """3x2 labelled grid: input|albedo|shading / depth|normal|recon.""" maps = extract_maps(output, img_rgb, ts) grid = Image.new("RGB", (ts * 3, ts * 2)) grid.paste(maps["input"], (0, 0)) grid.paste(maps["albedo"], (ts, 0)) grid.paste(maps["shading"], (ts * 2, 0)) grid.paste(maps["depth"], (0, ts)) grid.paste(maps["normal"], (ts, ts)) grid.paste(maps["recon"], (ts * 2, ts)) draw = ImageDraw.Draw(grid) for (x, y, name) in [(4, 4, "input"), (ts + 4, 4, "albedo"), (ts * 2 + 4, 4, "shading"), (4, ts + 4, "depth"), (ts + 4, ts + 4, "normal"), (ts * 2 + 4, ts + 4, "recon")]: draw.text((x, y), name, fill=(255, 255, 0)) return grid