ShadeNet-2-20M / inference_utils.py
singam96's picture
ShadeNet-2 20M: weights (EMA), ONNX fp32+int8, app + card
21e2ec3 verified
Raw History Blame Contribute Delete
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