File size: 6,427 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
"""Batched Bezier-stroke decomposer: fits strokes to many images at once.

Same stroke format and the same coarse-to-fine algorithm as stroke.py (11 numbers per stroke:
p0 p1 p2 control points, width, r g b, alpha), but every tensor carries a leading image dimension B.
Each image keeps its own strokes and its own loss. Adam is per-parameter, so images do not interact.

Memory notes: the per-stroke distance field is the big tensor. It is computed one curve segment at a time
and, with checkpoint=True, recomputed in the backward pass instead of stored.
"""
import torch
from torch.utils.checkpoint import checkpoint

N_PARAMS = 11


def grid(H, W, device):
    ys, xs = torch.meshgrid(
        (torch.arange(H, device=device) + 0.5) / H, (torch.arange(W, device=device) + 0.5) / W, indexing="ij"
    )
    return torch.stack([xs, ys], -1).reshape(-1, 2)  # (HW, 2)


def _min_dist2(pts, g):
    """pts (M, K+1, 2) polyline, g (HW, 2) -> squared distance from every pixel to the polyline, (M, HW)."""
    a, b = pts[:, :-1], pts[:, 1:]
    ab = b - a
    ab2 = (ab * ab).sum(-1) + 1e-8
    best = None
    for k in range(a.shape[1]):
        ak, abk = a[:, k, None], ab[:, k, None]  # (M, 1, 2)
        ag = g[None] - ak  # (M, HW, 2)
        u = ((ag * abk).sum(-1) / ab2[:, k, None]).clamp(0, 1)
        d = ag - u[..., None] * abk
        d2 = (d * d).sum(-1)
        best = d2 if best is None else torch.minimum(best, d2)
    return best


def stroke_masks(s, g, H, K=8, soft=0.75):
    """s (M, 11) -> soft coverage (M, HW)."""
    t = torch.linspace(0, 1, K + 1, device=s.device)[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
    d = (_min_dist2(pts, g) + 1e-12).sqrt()
    return torch.sigmoid((s[:, 6:7] / 2 - d) * H / soft)


def composite(s, masks, base):
    """s (B, N, 11), masks (B, N, HW), base (B, 3, HW). Paint strokes in order, closed-form 'over' compositing."""
    am = s[:, :, 10:11] * masks
    rev = torch.cumprod((1 - am).flip(1), 1).flip(1)  # rev[:, i] = prod_{j>=i} (1 - am_j)
    after = torch.cat([rev[:, 1:], torch.ones_like(rev[:, :1])], 1)  # strokes painted after i
    return torch.einsum("bnc,bnp->bcp", s[:, :, 7:10], am * after) + base * rev[:, 0:1]


def render(s, H, W, base=None, mask_fn=None, use_checkpoint=False):
    """s (B, N, 11) -> (B, 3, HW). base defaults to a white canvas."""
    B, N, _ = s.shape
    if base is None:
        base = torch.ones(B, 3, H * W, device=s.device)
    if N == 0:
        return base
    g = grid(H, W, s.device)
    fn = mask_fn or stroke_masks
    flat = s.reshape(B * N, N_PARAMS)
    if use_checkpoint:
        masks = checkpoint(fn, flat, g, H, use_reentrant=False)
    else:
        masks = fn(flat, g, H)
    return composite(s, masks.reshape(B, N, -1), base)


DEFAULT_CLAMP = {"pos": (-0.1, 1.1), "width": (0.004, 0.6), "alpha": (0.2, 1.0)}


def clamp_(s, pos=(-0.1, 1.1), width=(0.004, 0.6), alpha=(0.2, 1.0)):
    with torch.no_grad():
        s[..., 0:6].clamp_(*pos)
        s[..., 6].clamp_(*width)
        s[..., 7:10].clamp_(0, 1)
        s[..., 10].clamp_(*alpha)


def init_strokes(n, target, canvas, H, W, width):
    """Place n new strokes per image where that image's error is largest, colored from its target."""
    B = target.shape[0]
    dev = target.device
    err = (target - canvas).abs().sum(1) + 1e-6  # (B, HW)
    idx = torch.multinomial(err, n, replacement=err.shape[1] < n)  # (B, n)
    c = grid(H, W, dev)[idx]  # (B, n, 2)
    ang = torch.rand(B, n, device=dev) * 3.1416
    dirv = torch.stack([ang.cos(), ang.sin()], -1) * width * (0.5 + torch.rand(B, n, 1, device=dev))
    s = torch.empty(B, n, N_PARAMS, device=dev)
    s[..., 0:2] = c - dirv
    s[..., 2:4] = c + 0.1 * width * torch.randn(B, n, 2, device=dev)
    s[..., 4:6] = c + dirv
    s[..., 6] = width
    s[..., 7:10] = torch.gather(target, 2, idx[:, None, :].expand(-1, 3, -1)).transpose(1, 2)
    s[..., 10] = 0.9
    return s


def psnr(a, b):
    """Per-image PSNR, (B,)."""
    return 10 * torch.log10(1 / (a - b).pow(2).mean(dim=(1, 2)).clamp_min(1e-10))


def decompose(target, H, W, stages, steps=150, lr=0.01, mask_fn=None, use_checkpoint=False, clamp=None):
    """target (B, 3, HW) in [0, 1]. stages: [(n_strokes, init_width), ...] big brushes first.

    clamp: optional dict with keys pos/width/alpha (see DEFAULT_CLAMP). `steps` may be an int or one int per stage.
    Returns strokes (B, sum(n), 11), final canvas (B, 3, HW), and psnr_by_stage (B, len(stages)).
    """
    clamp = {**DEFAULT_CLAMP, **(clamp or {})}
    steps = [steps] * len(stages) if isinstance(steps, int) else list(steps)
    B = target.shape[0]
    dev = target.device
    frozen = torch.empty(B, 0, N_PARAMS, device=dev)
    canvas = torch.ones(B, 3, H * W, device=dev)
    stage_psnr = []
    for (n, width), n_steps in zip(stages, steps):
        s = init_strokes(n, target, canvas, H, W, width).requires_grad_(True)
        clamp_(s, **clamp)  # start inside the allowed ranges
        opt = torch.optim.Adam([s], lr=lr)
        for _ in range(n_steps):
            out = render(s, H, W, base=canvas, mask_fn=mask_fn, use_checkpoint=use_checkpoint)
            diff = out - target
            loss = (diff.pow(2).mean(dim=(1, 2)) + 0.5 * diff.abs().mean(dim=(1, 2))).sum()  # per-image, summed
            opt.zero_grad()
            loss.backward()
            opt.step()
            clamp_(s, **clamp)
        frozen = torch.cat([frozen, s.detach()], 1)
        with torch.no_grad():
            canvas = render(s.detach(), H, W, base=canvas, mask_fn=mask_fn)
        stage_psnr.append(psnr(canvas, target))
    return frozen, canvas, torch.stack(stage_psnr, 1)


def quantize(s, bins_pos=64, bins_w=16, bins_c=16, bins_a=4):
    """Snap every field to a bin center, to simulate the discrete tokenizer. Works on any leading shape."""
    q = s.clone()

    def snap(x, lo, hi, n):
        x = ((x - lo) / (hi - lo)).clamp(0, 1)
        return lo + (torch.round(x * (n - 1)) / (n - 1)) * (hi - lo)

    lo_w, hi_w = torch.log(torch.tensor(0.004)), torch.log(torch.tensor(0.6))
    q[..., 0:6] = snap(s[..., 0:6], -0.1, 1.1, bins_pos)
    q[..., 6] = torch.exp(snap(torch.log(s[..., 6]), lo_w.to(s.device), hi_w.to(s.device), bins_w))
    q[..., 7:10] = snap(s[..., 7:10], 0, 1, bins_c)
    q[..., 10] = snap(s[..., 10], 0.2, 1.0, bins_a)
    return q