Piccaso-0.1 / code /batched.py
shing-dev's picture
Piccaso-0.1: weights, config, inference code, model card
a431a1c verified
Raw History Blame Contribute Delete
6.43 kB
"""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