Download inference_utils.py from singam96/ShadeNet-3.2-5M: direct link, hf CLI and curl.
- Browser
- Download file 3.96 kB
-
https://huggingface.co/singam96/ShadeNet-3.2-5M/resolve/main/inference_utils.py
- Command line
-
hf download hf://singam96/ShadeNet-3.2-5M/inference_utils.py
-
curl -L -o inference_utils.py https://huggingface.co/singam96/ShadeNet-3.2-5M/resolve/main/inference_utils.py
3.96 kB
| """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 | |