File size: 15,695 Bytes
0ca5a80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1dfc73d
0ca5a80
 
 
 
 
1dfc73d
0ca5a80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1dfc73d
 
0ca5a80
1dfc73d
 
 
0ca5a80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
"""
Material maps for 3dvalley.com: one picture of a surface (a texture tile or a
photo) in, the maps a PBR renderer needs out, all tiling when the picture
tiles. One API endpoint for headless callers (the site's browser client),
plus a small demo UI.

How a run goes, all on the GPU in one call:
- Two small ESRGAN nets trained on texture sets (Joey Ballentine's Material
  Map Generator, Apache-2.0): one gives a tangent-space normal map, the other
  displacement and roughness. They are what reads a painted brick as raised
  and its white mortar as sunk, which brightness alone gets backwards.
- Lotus-G normal (Apache-2.0, Stable Diffusion 2 fine-tuned for surface
  normals, one step) gives the broad shape: the rounded top of a cobble, the
  bevel of a brick. Its low frequencies and the ESRGAN detail are merged in
  slope space with the slopes of the displacement map, so normal and height
  agree.
- CLIP (MIT) names the material class, which sets the roughness level and
  whether anything is metal; the ESRGAN roughness adds the variation.

Every convolution pads circularly (the ESRGAN input is wrapped, the Lotus UNet
and VAE have their padding mode switched), and every filter wraps, so a
seamless picture gives seamless maps.
"""

import os
import tempfile
import time

import spaces

os.environ["GRADIO_TEMP_DIR"] = os.path.join(tempfile.gettempdir(), "gradio")
os.makedirs(os.environ["GRADIO_TEMP_DIR"], exist_ok=True)

import gradio as gr
import numpy as np
import torch
import torch.nn.functional as F
from diffusers import AutoencoderKL, UNet2DConditionModel
from huggingface_hub import hf_hub_download
from PIL import Image
from spandrel import ModelLoader
from transformers import CLIPModel, CLIPProcessor, CLIPTextModel, CLIPTokenizer

from upsampler_theme import UPSAMPLER_CSS, UPSAMPLER_THEME, footer_html, header_html

DEVICE = "cuda"
MAX_SIDE = 1024
LOTUS_SIDE = 768  # Lotus-G is Stable Diffusion 2 base: 512 to 768 is home.
MAPS_REPO = "InvokeAI/pbr-material-maps"
MAPS_REVISION = "b7ca9ebc6e14688a69d41872d2b9c80ea453e8f0"
LOTUS_REPO = "jingheya/lotus-normal-g-v1-1"
CLIP_REPO = "openai/clip-vit-base-patch32"


def _log(*parts):
    print("[maps]", *parts, flush=True)


def _circular(module: torch.nn.Module) -> None:
    for m in module.modules():
        if isinstance(m, torch.nn.Conv2d) and m.padding not in (0, (0, 0)):
            m.padding_mode = "circular"


def _esrgan(name: str) -> torch.nn.Module:
    path = hf_hub_download(MAPS_REPO, name, revision=MAPS_REVISION)
    return ModelLoader().load_from_file(path).model.eval().half().to(DEVICE)


normal_net = _esrgan("normal_map_generator.safetensors")
franken_net = _esrgan("franken_map_generator.safetensors")

lotus_unet = UNet2DConditionModel.from_pretrained(LOTUS_REPO, subfolder="unet", torch_dtype=torch.float16).to(DEVICE)
lotus_vae = AutoencoderKL.from_pretrained(LOTUS_REPO, subfolder="vae", torch_dtype=torch.float16).to(DEVICE)
_circular(lotus_unet)
_circular(lotus_vae)
# Lotus runs with an empty prompt: encode it once and drop the text encoder.
with torch.no_grad():
    _tok = CLIPTokenizer.from_pretrained(LOTUS_REPO, subfolder="tokenizer")
    _enc = CLIPTextModel.from_pretrained(LOTUS_REPO, subfolder="text_encoder")
    _ids = _tok([""], padding="max_length", max_length=_tok.model_max_length, return_tensors="pt").input_ids
    EMPTY_PROMPT = _enc(_ids)[0].half().to(DEVICE)  # computed on the CPU, like the CLIP classes below
    del _tok, _enc
# The task embedding that selects the normal head (see Lotus's infer.py).
_task = torch.tensor([[1.0, 0.0]])
TASK_EMB = torch.cat([torch.sin(_task), torch.cos(_task)], dim=-1).half().to(DEVICE)

clip_model = CLIPModel.from_pretrained(CLIP_REPO).eval()
clip_processor = CLIPProcessor.from_pretrained(CLIP_REPO)

# (label, words for CLIP, roughness level, metal): "metal" is bare metal all
# over, "rust" is metal only where grey steel shows through.
CLASSES = [
    ("brick", "a brick wall texture", 0.85, None),
    ("stone", "a cobblestone or stone paving texture", 0.8, None),
    ("rock", "a rough natural rock texture", 0.85, None),
    ("concrete", "a concrete or plaster wall texture", 0.9, None),
    ("asphalt", "an asphalt road texture", 0.9, None),
    ("wood", "a wooden planks texture", 0.7, None),
    ("varnished wood", "a varnished polished wood floor texture", 0.35, None),
    ("bark", "a tree bark texture", 0.9, None),
    ("ground", "a dirt, soil, mud or sand ground texture", 0.95, None),
    ("vegetation", "a grass, moss or leaves texture", 0.8, None),
    ("marble", "a polished marble texture", 0.2, None),
    ("tiles", "a glazed ceramic tiles texture", 0.25, None),
    ("fabric", "a fabric, cloth or carpet texture", 0.9, None),
    ("leather", "a leather texture", 0.6, None),
    ("plastic", "a plastic surface texture", 0.4, None),
    ("painted metal", "a painted metal surface texture", 0.5, None),
    ("rusted metal", "a rusty corroded metal texture", 0.8, "rust"),
    ("brushed metal", "a brushed steel or aluminium metal texture", 0.35, "metal"),
    ("polished metal", "a shiny polished metal, chrome, gold or copper texture", 0.15, "metal"),
    ("snow", "a snow or ice texture", 0.3, None),
]
# The class words are embedded once, on the CPU: ZeroGPU only lends a GPU
# inside a @spaces.GPU call, so nothing runs on "cuda" at start-up.
with torch.no_grad():
    _t = clip_processor(text=[c[1] for c in CLASSES], return_tensors="pt", padding=True)
    CLASS_EMB = F.normalize(clip_model.get_text_features(**_t).float(), dim=-1).to(DEVICE)
clip_model = clip_model.half().to(DEVICE)
_log("models ready")


# --- plain-array helpers, all wrapping at the edges -------------------------

def _blur(a: torch.Tensor, sigma: float) -> torch.Tensor:
    """Separable Gaussian on an (H, W) tensor, wrapping around the edges."""
    if sigma <= 0:
        return a
    radius = max(1, int(3 * sigma))
    x = torch.arange(-radius, radius + 1, device=a.device, dtype=a.dtype)
    k = torch.exp(-(x**2) / (2 * sigma**2))
    k = k / k.sum()
    out = F.pad(a[None, None], (radius, radius, 0, 0), mode="circular")
    out = F.conv2d(out, k.view(1, 1, 1, -1))
    out = F.pad(out, (0, 0, radius, radius), mode="circular")
    return F.conv2d(out, k.view(1, 1, -1, 1))[0, 0]


def _resize_wrap(x: torch.Tensor, size: tuple[int, int]) -> torch.Tensor:
    """Resize (N, C, H, W) so the result still tiles: pad by wrapping, scale, crop."""
    h, w = x.shape[-2:]
    if (h, w) == size:
        return x
    pad = 4
    big = F.pad(x, (pad, pad, pad, pad), mode="circular")
    sy, sx = size[0] / h, size[1] / w
    out = F.interpolate(big, size=(round((h + 2 * pad) * sy), round((w + 2 * pad) * sx)), mode="bicubic", align_corners=False)
    oy, ox = round(pad * sy), round(pad * sx)
    return out[..., oy : oy + size[0], ox : ox + size[1]]


def _slopes(n: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """(3, H, W) normals, x right, y up → slopes dh/dx and dh/dy_up, tilt removed.
    A normal is (-dh/dx, -dh/dy, 1) normalised."""
    nz = n[2].clamp(min=0.2)
    p, q = -n[0] / nz, -n[1] / nz
    return p - p.mean(), q - q.mean()


def _grad(h: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    """Central differences with wrap: dh/dx and dh/dy_up (rows run down)."""
    gx = (torch.roll(h, -1, 1) - torch.roll(h, 1, 1)) / 2
    gy = (torch.roll(h, 1, 0) - torch.roll(h, -1, 0)) / 2
    return gx, gy


def _stretch(a: torch.Tensor, lo: float = 0.005, hi: float = 0.995) -> torch.Tensor:
    flat = a.flatten()
    if flat.numel() > 1_000_000:
        flat = flat[:: flat.numel() // 1_000_000 + 1]
    a_lo, a_hi = torch.quantile(flat, lo), torch.quantile(flat, hi)
    return ((a - a_lo) / (a_hi - a_lo).clamp(min=1e-6)).clamp(0, 1)


def _smoothstep(e0: float, e1: float, x: torch.Tensor) -> torch.Tensor:
    t = ((x - e0) / (e1 - e0)).clamp(0, 1)
    return t * t * (3 - 2 * t)


# --- the models --------------------------------------------------------------

def _run_esrgan(net: torch.nn.Module, rgb: torch.Tensor) -> torch.Tensor:
    """(1, 3, H, W) in [0, 1] → (3, H, W) in [0, 1]; wrapped so the edges tile."""
    pad = 32
    x = F.pad(rgb, (pad, pad, pad, pad), mode="circular").half()
    return net(x)[0, :, pad:-pad, pad:-pad].float().clamp(0, 1)


def _run_lotus(rgb: torch.Tensor) -> torch.Tensor:
    """(1, 3, H, W) in [0, 1] → (3, H, W) unit normals, x right, y up, z out."""
    h, w = rgb.shape[-2:]
    scale = min(1.0, LOTUS_SIDE / max(h, w))
    size = (max(64, round(h * scale / 64) * 64), max(64, round(w * scale / 64) * 64))
    x = _resize_wrap(rgb, size) * 2 - 1
    latents = lotus_vae.encode(x.half()).latent_dist.mode() * lotus_vae.config.scaling_factor
    noise = torch.randn(latents.shape, generator=torch.Generator(DEVICE).manual_seed(0), device=DEVICE, dtype=latents.dtype)
    x0 = lotus_unet(
        torch.cat([latents, noise], dim=1),
        torch.tensor([999], device=DEVICE),
        encoder_hidden_states=EMPTY_PROMPT,
        class_labels=TASK_EMB,
        return_dict=False,
    )[0]
    decoded = lotus_vae.decode(x0 / lotus_vae.config.scaling_factor, return_dict=False)[0].float().clamp(-1, 1)
    n = _resize_wrap(decoded, (h, w))[0]
    return n / n.norm(dim=0, keepdim=True).clamp(min=1e-6)


def _classify(image: Image.Image) -> tuple[torch.Tensor, list[tuple[str, float]]]:
    inputs = clip_processor(images=image, return_tensors="pt").to(DEVICE)
    emb = F.normalize(clip_model.get_image_features(pixel_values=inputs.pixel_values.half()).float(), dim=-1)
    probs = (100 * emb @ CLASS_EMB.T).softmax(dim=-1)[0]
    order = probs.argsort(descending=True)[:3].tolist()
    return probs, [(CLASSES[i][0], round(float(probs[i]), 3)) for i in order]


@spaces.GPU(duration=20)
@torch.no_grad()
def _maps(image: Image.Image):
    t0 = time.time()
    w, h = image.size
    rgb = torch.from_numpy(np.asarray(image, np.float32) / 255).permute(2, 0, 1)[None].to(DEVICE)
    s = max(w, h) / 768  # filter sizes were tuned at 768 px

    es_normal = _run_esrgan(normal_net, rgb) * 2 - 1
    franken = _run_esrgan(franken_net, rgb)
    lotus = _run_lotus(rgb)
    probs, top = _classify(image)
    _log(f"models {time.time() - t0:.2f}s", top)

    # Height: the texture-trained displacement. Normal: the displacement's
    # slopes, plus Lotus's broad shape and the ESRGAN normal's fine detail.
    height = _stretch(franken[2])
    pd, qd = _grad(height)
    pd, qd = pd * max(w, h) / 40, qd * max(w, h) / 40
    pl, ql = _slopes(lotus)
    pe, qe = _slopes(es_normal)
    p = (_blur(pl, 2 * s) + pe - _blur(pe, 3 * s) + pd) / 2
    q = (_blur(ql, 2 * s) + qe - _blur(qe, 3 * s) + qd) / 2
    normal = torch.stack([-p, -q, torch.ones_like(p)])
    normal = normal / normal.norm(dim=0, keepdim=True)

    # Roughness: the class sets the level, the ESRGAN map the variation.
    level = sum(float(probs[i]) * c[2] for i, c in enumerate(CLASSES))
    rough = franken[1]
    roughness = (level + (rough - rough.median()) * 1.2).clamp(0.04, 1)

    # Metallic: bare metal is metal all over; rusted metal only where grey
    # steel shows (low saturation). Everything else is not metal.
    metal = sum(float(probs[i]) for i, c in enumerate(CLASSES) if c[3] == "metal")
    rust = sum(float(probs[i]) for i, c in enumerate(CLASSES) if c[3] == "rust")
    mx, mn = rgb[0].max(dim=0).values, rgb[0].min(dim=0).values
    saturation = (mx - mn) / mx.clamp(min=1e-3)
    bare = _smoothstep(0.35, 0.15, saturation)
    metallic = _smoothstep(0.35, 0.65, metal + rust * bare)
    roughness = roughness - metallic * 0.15

    info = {
        "material": top[0][0],
        "classes": [{"label": label, "p": p_} for label, p_ in top],
        "roughness_level": round(level, 3),
        "metal": round(metal, 3),
        "gpu_seconds": round(time.time() - t0, 2),
    }
    out = (
        (normal.permute(1, 2, 0) * 0.5 + 0.5).clamp(0, 1).cpu().numpy(),
        height.cpu().numpy(),
        roughness.clamp(0, 1).cpu().numpy(),
        metallic.clamp(0, 1).cpu().numpy(),
    )
    torch.cuda.empty_cache()
    return out, info


def _save_png(array: np.ndarray, stem: str, bits: int = 8) -> str:
    path = os.path.join(os.environ["GRADIO_TEMP_DIR"], f"{stem}-{time.time_ns()}.png")
    if bits == 16:
        Image.fromarray((array * 65535).round().astype(np.uint16)).save(path)
    else:
        Image.fromarray((array * 255).round().astype(np.uint8)).save(path, optimize=False, compress_level=6)
    return path


def material_maps(image, directx: bool = False):
    """A picture of a surface → normal (OpenGL unless `directx`), height
    (16-bit), roughness and metallic PNGs at its size (capped at 1024 px),
    and a small JSON note of what the surface was taken for."""
    if image is None:
        raise gr.Error("Upload a picture of a surface.")
    if not isinstance(image, Image.Image):
        image = Image.open(image)
    image = image.convert("RGB")
    if max(image.size) > MAX_SIDE:
        scale = MAX_SIDE / max(image.size)
        image = image.resize((max(8, round(image.width * scale)), max(8, round(image.height * scale))), Image.LANCZOS)
    t0 = time.time()
    (normal, height, roughness, metallic), info = _maps(image)
    if directx:
        normal = normal.copy()
        normal[..., 1] = 1 - normal[..., 1]
    info["convention"] = "directx" if directx else "opengl"
    info["size"] = [image.width, image.height]
    files = (
        _save_png(normal, "normal"),
        _save_png(height, "height", bits=16),
        _save_png(roughness, "roughness"),
        _save_png(metallic, "metallic"),
    )
    info["seconds"] = round(time.time() - t0, 2)
    _log("done", info)
    return (*files, info)


with gr.Blocks(title="Material Maps - Normal, Height and Roughness from One Image") as demo:
    gr.HTML(header_html(
        "Material Maps",
        "Normal, height, roughness and metallic maps from one picture of a surface. Seamless in, seamless out.",
    ))
    with gr.Row(equal_height=False):
        with gr.Column():
            src = gr.Image(type="pil", image_mode="RGB", label="Texture or photo of a surface", height=360)
            directx = gr.Checkbox(value=False, label="DirectX normal map (green down, for Unreal)")
            btn = gr.Button("Make Maps", variant="primary")
        with gr.Column():
            with gr.Row():
                out_normal = gr.Image(type="filepath", label="Normal", height=200)
                out_height = gr.Image(type="filepath", label="Height (16-bit)", height=200)
            with gr.Row():
                out_rough = gr.Image(type="filepath", label="Roughness", height=200)
                out_metal = gr.Image(type="filepath", label="Metallic", height=200)
            out_info = gr.JSON(label="Surface")
    btn.click(
        material_maps,
        inputs=[src, directx],
        outputs=[out_normal, out_height, out_rough, out_metal, out_info],
        api_name="material_maps",
    )
    gr.HTML(footer_html(
        "Turn a texture or a photo of a surface into a PBR material: a tangent-space normal map, a 16-bit "
        "height (displacement) map, roughness and metallic, at the picture's size up to 1024 pixels. Nets "
        "trained on texture sets read painted bricks and stones the right way round, a surface-normal "
        "diffusion model adds the broad shape, and every step wraps at the edges so seamless textures stay "
        "seamless. Ready for Blender, Unity, Unreal, three.js and glTF.",
        "https://upsampler.com",
        "Upsampler",
    ))

if __name__ == "__main__":
    demo.queue(default_concurrency_limit=2).launch(
        theme=UPSAMPLER_THEME, css=UPSAMPLER_CSS, ssr_mode=False, show_error=True
    )