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