spritegen-cli 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
spritegen/upscale.py ADDED
@@ -0,0 +1,444 @@
1
+ """Real-ESRGAN, locally and on the GPU, so a paid call never has to buy resolution.
2
+
3
+ The expensive calls here are billed by pixels and by seconds. Generating at the cheap
4
+ tier and enlarging afterwards moves the quality decision off the invoice: a `low` anchor
5
+ and a 480p clip cost a fraction of the same thing at the top tier, and what they give up
6
+ is resolution — which is the one thing a model like this puts back.
7
+
8
+ **Ported from an `upscale_atlas.py` that reimplements RRDBNet without `basicsr`.** The
9
+ weights are downloaded rather than vendored: the ncnn `.bin`/`.param` files that ship
10
+ with Real-ESRGAN are for a different runtime, and its `realesrgan-ncnn-vulkan` binary is
11
+ a Linux ELF that does not run everywhere this does.
12
+
13
+ What changed in the port is the alpha, and it matters more here than it did there. A
14
+ matted board is transparent over most of its area, and the RGB behind that transparency
15
+ is whatever the matte left — usually black. Handing that to the network sharpens the
16
+ black into the figure's edge, and every frame comes back with a dark rim. So the
17
+ transparent region is filled from the nearest opaque colour before the network sees it,
18
+ and the alpha channel is enlarged on its own.
19
+ """
20
+
21
+ from __future__ import annotations
22
+
23
+ import argparse
24
+ import hashlib
25
+ from pathlib import Path
26
+
27
+ from . import fal
28
+ from .fal import Untrusted
29
+
30
+ #: The two the project already uses. `anime` is the 6-block model and is the one for
31
+ #: character art; `general` is the full 23-block model, for a photographic reference.
32
+ MODEL_URLS = {
33
+ "anime": (
34
+ "https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.2.4/"
35
+ "RealESRGAN_x4plus_anime_6B.pth"
36
+ ),
37
+ "general": (
38
+ "https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth"
39
+ ),
40
+ }
41
+
42
+ #: Where a weight file may come from. The URLs above are constants, so this is
43
+ #: defence in depth rather than a live path — and the line that has to change first if
44
+ #: they ever stop being constants.
45
+ WEIGHT_HOSTS = ("github.com", "objects.githubusercontent.com")
46
+
47
+ #: The network is 4x whatever it is given; any other factor is a resize afterwards.
48
+ NATIVE_SCALE = 4
49
+
50
+ #: Inference is tiled, so a 3732x3320 board does not have to fit in VRAM at once.
51
+ TILE = 256
52
+ TILE_PAD = 10
53
+
54
+ #: How far the opaque colour is flooded outward before the network sees the image.
55
+ FILL_PASSES = 8
56
+
57
+ INSTALL = (
58
+ "the upscale needs torch, which is an optional extra: "
59
+ 'pip install "spritegen-cli[upscale]", or uv sync --extra upscale'
60
+ )
61
+
62
+
63
+ def cache_dir() -> Path:
64
+ """Where a weight file lives — `SPRITEGEN_CACHE`, or the default beside it.
65
+
66
+ Outside any workspace on purpose. A weight file is tens of megabytes, it is the same
67
+ file for every project this CLI is run in, and it can always be fetched again — none
68
+ of which describes something that belongs in somebody's repository.
69
+ """
70
+ from . import settings
71
+
72
+ return settings.load().cache / "weights"
73
+
74
+
75
+ def require_torch():
76
+ """`torch`, or an error naming the one command that fixes it.
77
+
78
+ Resolved here and never at module scope. Torch is some two gigabytes with CUDA, and
79
+ this is one free command of a CLI whose other commands have to keep answering on a
80
+ machine that never installed it — `status` and `cost` especially.
81
+ """
82
+ try:
83
+ import torch
84
+ except ModuleNotFoundError as exc:
85
+ raise NotImplementedError(INSTALL) from exc
86
+ return torch
87
+
88
+
89
+ def available() -> bool:
90
+ """Whether the upscale could run at all. Used to report, never to decide silently."""
91
+ try:
92
+ require_torch()
93
+ except NotImplementedError:
94
+ return False
95
+ return True
96
+
97
+
98
+ def pick_device(prefer: str = "auto"):
99
+ """CUDA where there is one, the CPU where there is not — or whatever was asked for.
100
+
101
+ Reported rather than assumed: the difference on a board is seconds against minutes,
102
+ and a run that quietly fell back to the CPU looks like a hang. Asking for `cuda`
103
+ where there is none is an error rather than a silent fallback, for the same reason.
104
+ """
105
+ torch = require_torch()
106
+ if prefer == "cpu":
107
+ return torch.device("cpu")
108
+ if prefer == "cuda":
109
+ if not torch.cuda.is_available():
110
+ raise NotImplementedError(
111
+ "cuda was asked for and torch reports none; leave the device on auto "
112
+ "for the CPU, or install a torch built against this machine's CUDA"
113
+ )
114
+ return torch.device("cuda")
115
+ return torch.device("cuda" if torch.cuda.is_available() else "cpu")
116
+
117
+
118
+ def download_weights(model: str) -> Path:
119
+ """The weight file, out of the cache or out of the release that holds it.
120
+
121
+ Held to the same discipline `fal.download` is, and for the same reason — it is the
122
+ other place in this tool that writes bytes from the network to disk:
123
+
124
+ - **HTTPS to the host the constant names**, checked rather than assumed. The URLs are
125
+ constants, so this is not an attacker's path today; it is what stops one becoming
126
+ an attacker's path the first time somebody makes them configurable.
127
+ - **A size ceiling**, so a redirect to something enormous fails instead of filling
128
+ the disk.
129
+ - **A digest written beside the file.** Nothing upstream publishes a hash to pin
130
+ against, so this cannot say the *first* download was the right one. What it can say
131
+ is that the file has not changed since — which is the case that matters for a cache
132
+ trusted forever and never looked at again.
133
+
134
+ `torch.load(..., weights_only=True)` in `load_model` is what keeps a swapped file
135
+ from being code execution; this keeps it from being silent.
136
+ """
137
+ if model not in MODEL_URLS:
138
+ known = ", ".join(sorted(MODEL_URLS))
139
+ raise KeyError(f"unknown upscale model {model!r}; there are: {known}")
140
+
141
+ cache = cache_dir()
142
+ cache.mkdir(parents=True, exist_ok=True)
143
+ destination = cache / f"{model}.pth"
144
+ digest_at = destination.with_suffix(".sha256")
145
+
146
+ if destination.exists():
147
+ if digest_at.is_file():
148
+ recorded = digest_at.read_text(encoding="utf-8").strip()
149
+ found = _digest(destination)
150
+ if found != recorded:
151
+ raise Untrusted(
152
+ f"{destination} has changed since it was downloaded "
153
+ f"({found[:12]} against {recorded[:12]}); delete it to fetch again"
154
+ )
155
+ return destination
156
+
157
+ _fetch(MODEL_URLS[model], destination)
158
+ digest_at.write_text(_digest(destination) + "\n", encoding="utf-8")
159
+ return destination
160
+
161
+
162
+ def _digest(path: Path) -> str:
163
+ found = hashlib.sha256()
164
+ with path.open("rb") as handle:
165
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
166
+ found.update(chunk)
167
+ return found.hexdigest()
168
+
169
+
170
+ def _fetch(url: str, destination: Path) -> Path:
171
+ """One weight file, through the same fetch a paid result comes through.
172
+
173
+ `fal.download` with a different host set, and not a second implementation of it. The
174
+ first version here checked the scheme and the host of the URL it was *given* — which
175
+ is not the host it ends up talking to, because `urlopen` follows redirects inside the
176
+ library and a GitHub release URL redirects by design. `download` checks every hop and
177
+ where the chain lands, which is what that has to mean.
178
+
179
+ Written to a `.partial` first: a fetch interrupted halfway would otherwise leave a
180
+ truncated file that `exists()` accepts forever.
181
+ """
182
+ partial = destination.with_suffix(".partial")
183
+ try:
184
+ fal.download(url, partial, WEIGHT_HOSTS)
185
+ partial.replace(destination)
186
+ finally:
187
+ partial.unlink(missing_ok=True)
188
+ return destination
189
+
190
+
191
+ def fill_transparent(image):
192
+ """Take the black out from behind the figure before the network sees it.
193
+
194
+ The network sees RGB and no alpha. Where the matte left black behind the
195
+ transparency, that black is a hard edge as far as the network is concerned and it
196
+ sharpens it into the figure — which is the dark rim around an enlarged sprite.
197
+
198
+ Two treatments, because the transparent area is two different things. The band just
199
+ outside the silhouette is where the network actually looks, and it is inpainted from
200
+ the art beside it, so the edge continues instead of stopping. Everything beyond that
201
+ band is out of reach of any kernel and only has to not be black, so it takes the
202
+ figure's own mean colour.
203
+
204
+ An image with no transparency comes back converted and otherwise untouched, which is
205
+ the anchor's case before it has been cut.
206
+ """
207
+ import cv2
208
+ import numpy as np
209
+ from PIL import Image
210
+
211
+ rgba = np.asarray(image.convert("RGBA"))
212
+ opaque = rgba[:, :, 3] > 0
213
+ if not opaque.any() or opaque.all():
214
+ return Image.fromarray(np.ascontiguousarray(rgba[:, :, :3]), "RGB")
215
+
216
+ rgb = np.ascontiguousarray(rgba[:, :, :3])
217
+ rgb[~opaque] = rgb[opaque].mean(axis=0).round().astype(np.uint8)
218
+
219
+ solid = opaque.astype(np.uint8)
220
+ kernel = np.ones((3, 3), np.uint8)
221
+ reach = cv2.dilate(solid, kernel, iterations=FILL_PASSES)
222
+ band = ((reach > 0) & ~opaque).astype(np.uint8)
223
+ if band.any():
224
+ rgb = cv2.inpaint(rgb, band, FILL_PASSES, cv2.INPAINT_TELEA)
225
+ return Image.fromarray(rgb, "RGB")
226
+
227
+
228
+ def load_model(model: str, device):
229
+ """The network with its weights in it, on the device that will run it."""
230
+ torch = require_torch()
231
+ from .rrdb import RRDBNet
232
+
233
+ net = RRDBNet(in_nc=3, out_nc=3, nf=64, nb=6 if model == "anime" else 23, gc=32)
234
+ state = torch.load(str(download_weights(model)), map_location=device, weights_only=True)
235
+ # A Real-ESRGAN checkpoint wraps its weights; which key depends on the release.
236
+ for key in ("params_ema", "params"):
237
+ if key in state:
238
+ state = state[key]
239
+ break
240
+ net.load_state_dict(state, strict=True)
241
+ net.eval()
242
+ net.to(device)
243
+ return net
244
+
245
+
246
+ def tiles(height: int, width: int, tile: int = TILE, pad: int = TILE_PAD) -> list[dict]:
247
+ """How a `height x width` image is cut into overlapping tiles and put back together.
248
+
249
+ Pure arithmetic, and separate from the forward pass on purpose: this is the part that
250
+ is easy to get wrong, the way it goes wrong is a visible seam every `tile` pixels,
251
+ and a seam is not something a test of the network would notice. Extracted, it can be
252
+ checked against what it has to be — that the kept regions tile the output exactly,
253
+ with no gap and no overlap.
254
+
255
+ Each entry has three boxes, in input pixels for `read` and output pixels for the
256
+ other two:
257
+
258
+ - `read` — what goes into the network, the tile plus `pad` of context on each side
259
+ that exists. The context is what stops the network seeing a hard edge where the
260
+ tile was cut, which is the seam.
261
+ - `take` — where this tile's own pixels sit inside what came back, which is the pad
262
+ offset scaled up.
263
+ - `put` — where they go in the finished image.
264
+
265
+ `take` and `put` are the same size by construction; the pad is read and thrown away.
266
+ """
267
+ found = []
268
+ for top in range(0, height, tile):
269
+ for left in range(0, width, tile):
270
+ y0, x0 = max(top - pad, 0), max(left - pad, 0)
271
+ y1, x1 = min(top + tile + pad, height), min(left + tile + pad, width)
272
+ keep_h = min(tile, height - top) * NATIVE_SCALE
273
+ keep_w = min(tile, width - left) * NATIVE_SCALE
274
+ found.append(
275
+ {
276
+ "read": (y0, y1, x0, x1),
277
+ "take": (
278
+ (top - y0) * NATIVE_SCALE,
279
+ (top - y0) * NATIVE_SCALE + keep_h,
280
+ (left - x0) * NATIVE_SCALE,
281
+ (left - x0) * NATIVE_SCALE + keep_w,
282
+ ),
283
+ "put": (
284
+ top * NATIVE_SCALE,
285
+ top * NATIVE_SCALE + keep_h,
286
+ left * NATIVE_SCALE,
287
+ left * NATIVE_SCALE + keep_w,
288
+ ),
289
+ }
290
+ )
291
+ return found
292
+
293
+
294
+ def _tiled(net, tensor):
295
+ """One forward pass, in tiles, so a large board need not fit in VRAM at once."""
296
+ torch = require_torch()
297
+ _, _, height, width = tensor.shape
298
+ if height <= TILE and width <= TILE:
299
+ return net(tensor)
300
+
301
+ output = torch.zeros((1, 3, height * NATIVE_SCALE, width * NATIVE_SCALE), device=tensor.device)
302
+ for one in tiles(height, width):
303
+ ry0, ry1, rx0, rx1 = one["read"]
304
+ ty0, ty1, tx0, tx1 = one["take"]
305
+ py0, py1, px0, px1 = one["put"]
306
+ done = net(tensor[:, :, ry0:ry1, rx0:rx1])
307
+ output[:, :, py0:py1, px0:px1] = done[:, :, ty0:ty1, tx0:tx1]
308
+ return output
309
+
310
+
311
+ def enlarge(image, scale: int, model: str | None = None, device=None):
312
+ """One image, `scale` times larger, with the alpha kept out of the network.
313
+
314
+ The network is 4x and only 4x, so any other factor is a resize of its output. Going
315
+ through it is still worth it at 2x: the result carries detail the model invented,
316
+ which is the entire reason to run it rather than resize the input.
317
+ """
318
+ # Before torch is resolved: a scale under two is the caller's mistake either way,
319
+ # and being told about torch instead would send them to fix the wrong thing.
320
+ if scale < 2:
321
+ raise ValueError(f"an upscale is 2 or more; got {scale}")
322
+
323
+ torch = require_torch()
324
+ import numpy as np
325
+ from PIL import Image
326
+
327
+ from . import settings
328
+
329
+ config = settings.load()
330
+ model = model or config.upscale_model
331
+ device = pick_device(config.upscale_device) if device is None else device
332
+ net = load_model(model, device)
333
+
334
+ array = np.asarray(fill_transparent(image), dtype=np.float32) / 255.0
335
+ tensor = torch.from_numpy(array).permute(2, 0, 1).unsqueeze(0).to(device)
336
+ with torch.no_grad():
337
+ out = _tiled(net, tensor).squeeze(0).clamp(0, 1).cpu()
338
+
339
+ enlarged = Image.fromarray((out.permute(1, 2, 0).numpy() * 255).round().astype(np.uint8), "RGB")
340
+ target = (image.size[0] * scale, image.size[1] * scale)
341
+ if enlarged.size != target:
342
+ enlarged = enlarged.resize(target, Image.LANCZOS)
343
+
344
+ if image.mode != "RGBA":
345
+ return enlarged
346
+ result = enlarged.convert("RGBA")
347
+ result.putalpha(image.getchannel("A").resize(target, Image.LANCZOS))
348
+ return result
349
+
350
+
351
+ def enlarge_file(path: Path, scale: int, model: str | None = None, device=None) -> dict:
352
+ """Enlarge one file in place, keeping the original beside it the first time.
353
+
354
+ `<name>.raw.png` is written once and never overwritten, so running the upscale a
355
+ second time cannot lose what came back from the paid call — and the second run
356
+ starts from that original rather than from an already-enlarged image, which would
357
+ compound the model's invention twice over.
358
+
359
+ That backup is also why this command needs no ledger line: nothing it does costs
360
+ anything, and nothing it does is unrecoverable.
361
+ """
362
+ from PIL import Image
363
+
364
+ raw = path.with_name(f"{path.stem}.raw{path.suffix}")
365
+ if not raw.exists():
366
+ raw.write_bytes(path.read_bytes())
367
+ image = Image.open(raw).convert("RGBA")
368
+
369
+ from . import settings
370
+
371
+ model = model or settings.load().upscale_model
372
+ result = enlarge(image, scale, model, device)
373
+ result.save(path)
374
+
375
+ return {
376
+ "file": path.name,
377
+ "raw": raw.name,
378
+ "from": f"{image.size[0]}x{image.size[1]}",
379
+ "to": f"{result.size[0]}x{result.size[1]}",
380
+ "scale": scale,
381
+ "model": model,
382
+ }
383
+
384
+
385
+ def cmd_upscale(args: argparse.Namespace) -> int:
386
+ """Enlarge one artifact of one sprite, locally and for nothing.
387
+
388
+ Named by artifact rather than by stage, like everything else in the workspace: what
389
+ is being enlarged is the anchor, or the frames, not "whatever the video stage left".
390
+ Without one, it takes what the sprite produced most recently, which is almost always
391
+ what somebody means right after producing it.
392
+
393
+ No ledger line, because nothing was spent. The original is kept as `<name>.raw.png`
394
+ the first time and every later run starts from that — enlarging an already-enlarged
395
+ image would compound the model's invention twice over.
396
+
397
+ ASCII only, for the reason `workspace.cmd_status` is.
398
+ """
399
+ from . import workspace
400
+
401
+ state = workspace.load(workspace.check_name(args.name))
402
+ if state.outdated():
403
+ raise workspace.StageRefused(
404
+ f"{state.name}: this workspace is an older layout; run `spritegen migrate` first"
405
+ )
406
+
407
+ key = args.artifact or (state.last() and state.stages[state.last()].get("produced"))
408
+ if not key:
409
+ raise workspace.StageRefused(f"{state.name}: nothing produced yet, so nothing to enlarge")
410
+ entry = state.artifacts.get(key)
411
+ if entry is None:
412
+ known = ", ".join(state.artifacts) or "none"
413
+ raise workspace.StageRefused(f"{state.name}: no artifact {key!r}; it has: {known}")
414
+
415
+ directory = workspace.inside(state.name, entry["dir"])
416
+ targets = [
417
+ path
418
+ for path in sorted(directory.glob("*.png"))
419
+ if not path.stem.endswith(".raw") and not path.name.endswith(".raw.png")
420
+ ]
421
+ if not targets:
422
+ raise workspace.StageRefused(f"{state.name}: no image in {directory}")
423
+
424
+ if args.dry_run:
425
+ from . import settings
426
+
427
+ config = settings.load()
428
+ print(f"dry run: {args.scale}x {config.upscale_model} on {config.upscale_device}")
429
+ for path in targets:
430
+ print(f" {path}")
431
+ return 0
432
+
433
+ device = pick_device(args.device)
434
+ print(f"{args.scale}x on {device}")
435
+ for path in targets:
436
+ report = enlarge_file(path, args.scale, args.model, device)
437
+ print(f" {report['file']} {report['from']} -> {report['to']} (kept {report['raw']})")
438
+
439
+ # The files changed but nothing was produced, so `artifacts` is corrected rather than
440
+ # added to: `show` compares this against the disk, and leaving it stale would report
441
+ # drift that is this command's own doing.
442
+ entry["files"] = sorted(path.name for path in directory.glob("*") if path.is_file())
443
+ workspace.save(state)
444
+ return 0