Download code/detail_level.py from shing-dev/Piccaso-0.1: direct link, hf CLI and curl.
- Browser
- Download file 7.76 kB
-
https://huggingface.co/shing-dev/Piccaso-0.1/resolve/main/code/detail_level.py
- Command line
-
hf download hf://shing-dev/Piccaso-0.1/code/detail_level.py
-
curl -L -o detail_level.py https://huggingface.co/shing-dev/Piccaso-0.1/resolve/main/code/detail_level.py
7.76 kB
| """Detail level for extractor v3: a g x g grid of small strokes fitted at high resolution ON TOP of the base strokes. | |
| Why a special renderer: a plain render costs (strokes x all pixels); at 256 px with 196 strokes that is ~5x the whole base fit. | |
| Here every detail stroke is kept inside a box around its own cell (control points within +-`reach` cells of the cell | |
| centre, width <= one cell), so its paint can only land in a small window. Only window pixels are computed. | |
| Order: cells are visited in M*M = 16 groups (row % 4, col % 4), row-major inside a group. Windows of one group never overlap, | |
| so a whole group is composited at once and the result equals painting the strokes one by one in that order. | |
| Slot order is therefore group-major; it is fixed and identical for every image (the model sees the same slot->cell map). | |
| detail_order(g) -> list of cell ids in paint order | |
| fit_detail(target_hi, base_hi, g, steps) -> strokes (B, g*g, 11) in paint order, canvas (B,3,HW) | |
| """ | |
| import math | |
| import torch | |
| import batched | |
| REACH = 1.0 # control points may move this many cells away from the cell centre | |
| SOFT = 0.75 # same soft edge as batched.stroke_masks (pixels) | |
| MARGIN = 5 # pixels of soft edge kept inside a window (coverage there < 0.2%) | |
| M = 4 # cells of one group are M apart; windows (2*(REACH+0.5) cells + 2*MARGIN px) must fit in M cells | |
| def detail_order(g, m=M): | |
| return [r * g + c for gr in range(m) for gc in range(m) for r in range(gr, g, m) for c in range(gc, g, m)] | |
| def _windows(g, H, dev): | |
| """Per cell: flat pixel indices of its window in a padded canvas, and the window pixel centres in [0,1] coords.""" | |
| cell = H / g | |
| half = math.ceil(cell * (REACH + 0.5) + MARGIN) # reach + half the max width + soft edge | |
| assert 2 * half <= M * cell + 1, "windows of one group would overlap" | |
| P = 2 * half | |
| pad = half + 1 | |
| Hp = H + 2 * pad | |
| idx, ctr = [], [] | |
| oy, ox = torch.meshgrid(torch.arange(P, device=dev), torch.arange(P, device=dev), indexing="ij") | |
| for k in range(g * g): | |
| r, c = divmod(k, g) | |
| cy, cx = (r + 0.5) * cell, (c + 0.5) * cell | |
| y0, x0 = int(math.floor(cy)) - half, int(math.floor(cx)) - half | |
| yy, xx = y0 + oy, x0 + ox # unpadded pixel coords (may be outside the image) | |
| idx.append(((yy + pad) * Hp + (xx + pad)).reshape(-1)) | |
| ctr.append(torch.stack([(xx + 0.5) / H, (yy + 0.5) / H], -1).reshape(-1, 2).float()) | |
| return torch.stack(idx), torch.stack(ctr), pad, Hp # (g*g, P*P), (g*g, P*P, 2) | |
| def _pad(img, H, pad, Hp): | |
| B = img.shape[0] | |
| out = torch.ones(B, 3, Hp, Hp, device=img.device, dtype=img.dtype) | |
| out[:, :, pad:pad + H, pad:pad + H] = img.view(B, 3, H, H) | |
| return out.view(B, 3, -1) | |
| def _crop(imgp, H, pad, Hp): | |
| B = imgp.shape[0] | |
| return imgp.view(B, 3, Hp, Hp)[:, :, pad:pad + H, pad:pad + H].reshape(B, 3, H * H) | |
| def local_masks(s, ctr, H, K=4): | |
| """s (B, n, 11), ctr (n, P, 2) -> coverage (B, n, P) on each stroke's own window.""" | |
| B, n, _ = s.shape | |
| t = torch.linspace(0, 1, K + 1, device=s.device)[None, None, :, None] | |
| p0, p1, p2 = s[..., None, 0:2], s[..., None, 2:4], s[..., None, 4:6] | |
| pts = (1 - t) ** 2 * p0 + 2 * (1 - t) * t * p1 + t ** 2 * p2 # (B, n, K+1, 2) | |
| a, ab = pts[:, :, :-1], pts[:, :, 1:] - pts[:, :, :-1] | |
| ab2 = (ab * ab).sum(-1) + 1e-8 | |
| best = None | |
| for k in range(K): | |
| ag = ctr[None] - a[:, :, k, None] # (B, n, P, 2) | |
| u = ((ag * ab[:, :, k, None]).sum(-1) / ab2[:, :, k, None]).clamp(0, 1) | |
| d = ag - u[..., None] * ab[:, :, k, None] | |
| d2 = (d * d).sum(-1) | |
| best = d2 if best is None else torch.minimum(best, d2) | |
| return torch.sigmoid((s[..., 6:7] / 2 - (best + 1e-12).sqrt()) * H / SOFT) | |
| _lm_compiled = None | |
| def render_detail(s, base, g, H, win=None, compiled=False): | |
| """s (B, g*g, 11) in paint order (detail_order), base (B, 3, H*H) -> canvas (B, 3, H*H). Differentiable.""" | |
| idx, ctr, pad, Hp = win or _windows(g, H, s.device) | |
| order = detail_order(g) | |
| cv = _pad(base, H, pad, Hp) | |
| m = M | |
| per = [len(range(gr, g, m)) * len(range(gc, g, m)) for gr in range(m) for gc in range(m)] | |
| j = 0 | |
| for n in per: | |
| cells = torch.tensor(order[j:j + n], device=s.device) | |
| sg = s[:, j:j + n] | |
| lm = _lm_compiled if compiled and _lm_compiled is not None else local_masks | |
| am = sg[..., 10:11] * lm(sg, ctr[cells], H) # (B, n, P) | |
| ix = idx[cells].reshape(-1) # windows of one group are disjoint | |
| loc = cv[:, :, ix].view(cv.shape[0], 3, n, -1) | |
| new = loc * (1 - am[:, None]) + sg[..., 7:10].permute(0, 2, 1)[..., None] * am[:, None] | |
| cv = cv.index_copy(2, ix, new.reshape(cv.shape[0], 3, -1)) | |
| j += n | |
| return _crop(cv, H, pad, Hp) | |
| def clamp_detail(s, g, order_t): | |
| """Keep each stroke inside its box: control points within REACH cells of its cell centre, width <= one cell.""" | |
| with torch.no_grad(): | |
| r, c = order_t // g, order_t % g | |
| cen = torch.stack([(c + 0.5) / g, (r + 0.5) / g], -1) # (n, 2) x, y | |
| lo, hi = cen - REACH / g, cen + REACH / g | |
| for j in (0, 2, 4): | |
| s[..., j:j + 2] = torch.maximum(torch.minimum(s[..., j:j + 2], hi), lo) | |
| s[..., 6].clamp_(0.004, 1.0 / g) | |
| s[..., 7:10].clamp_(0, 1) | |
| s[..., 10] = 1.0 | |
| def init_detail(target, canvas, g, H): | |
| """Like extract_v2.init_anchored (error-weighted centroid + colour per cell), returned in paint order.""" | |
| from extract_v2 import init_anchored | |
| s = init_anchored(target, canvas, g, 0.6 / g, H) | |
| return s[:, detail_order(g)] | |
| def fit_detail(target, base, g, H, steps=100, lr=0.004, compiled=False): | |
| """target, base (B, 3, H*H) at the high resolution. Returns strokes in paint order and the final canvas.""" | |
| dev = target.device | |
| order_t = torch.tensor(detail_order(g), device=dev) | |
| win = _windows(g, H, dev) | |
| s = init_detail(target, base, g, H) | |
| clamp_detail(s, g, order_t) | |
| s.requires_grad_(True) | |
| opt = torch.optim.Adam([s], lr=lr) | |
| global _lm_compiled | |
| if compiled and _lm_compiled is None: | |
| _lm_compiled = torch.compile(local_masks, dynamic=False) | |
| for _ in range(steps): | |
| out = render_detail(s, base, g, H, win, compiled) | |
| diff = out - target | |
| loss = (diff.pow(2).mean(dim=(1, 2)) + 0.5 * diff.abs().mean(dim=(1, 2))).sum() | |
| opt.zero_grad() | |
| loss.backward() | |
| opt.step() | |
| clamp_detail(s, g, order_t) | |
| with torch.no_grad(): | |
| canvas = render_detail(s.detach(), base, g, H, win) | |
| return s.detach(), canvas | |
| def detail_keep(s, base, target, g, H, prune): | |
| """Keep a detail stroke only if it is visible (removal changes the picture by >= prune, extract_v2 units) AND it helps: | |
| removing it would make the error against the real image larger. Greedy, one pass in reverse paint order.""" | |
| win = _windows(g, H, s.device) | |
| s = s.clone() | |
| err = lambda img: (img - target).pow(2).mean((1, 2)) | |
| full = render_detail(s, base, g, H, win) | |
| e_full = err(full) | |
| keep = torch.ones(s.shape[:2], dtype=torch.bool, device=s.device) | |
| for i in reversed(range(s.shape[1])): | |
| a = s[:, i, 10].clone() | |
| s[:, i, 10] = 0 | |
| wo = render_detail(s, base, g, H, win) | |
| vis = (wo - full).abs().sum(1).mean(1) >= prune | |
| k = vis & (err(wo) > e_full) | |
| keep[:, i] = k | |
| s[:, i, 10] = torch.where(k, a, torch.zeros_like(a)) # dropped strokes stay off | |
| full = torch.where(k[:, None, None], full, wo) | |
| e_full = torch.where(k, e_full, err(wo)) | |
| return keep | |