File size: 7,764 Bytes
a431a1c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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


@torch.no_grad()
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