File size: 4,904 Bytes
181c7ca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""PXG-Tiny CLI (offline, NumPy-only).

  python3 cli.py "a golden sword" -o sword.png
  python3 cli.py "a golden sword" --variations 4 --sheet row.png
  python3 cli.py --ask "can you draw a photo of a cat"
  python3 cli.py --sheet-from-prompts prompts.txt --outdir sprites/
"""
import argparse
import json
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[0]
sys.path.insert(0, str(ROOT / "src"))

from PIL import Image  # noqa: E402

from pxg_tiny import config as C  # noqa: E402
from pxg_tiny.pipeline import PXGPipeline  # noqa: E402


def sheet(grids, path, scale=8):
    n = len(grids)
    im = Image.new("RGBA", (n * (16 * scale + 4) + 4, 16 * scale + 8),
                   (24, 24, 30, 255))
    from pxg_tiny.render import grid_to_rgba
    for i, g in enumerate(grids):
        tile = Image.fromarray(grid_to_rgba(g), "RGBA").resize(
            (16 * scale, 16 * scale), Image.NEAREST)
        im.paste(tile, (4 + i * (16 * scale + 4), 4), tile)
    im.save(path)


def main():
    ap = argparse.ArgumentParser(prog="pxg-tiny")
    ap.add_argument("prompt", nargs="?")
    ap.add_argument("-o", "--out", default=None)
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--temperature", type=float, default=None)
    ap.add_argument("--top-k", type=int, default=None)
    ap.add_argument("--variations", type=int, default=1)
    ap.add_argument("--sheet", action="store_true",
                    help="save a variations sheet next to the png")
    ap.add_argument("--sheet-from-prompts", default=None,
                    help="path to a text file with one prompt per line; "
                         "renders every prompt into --outdir")
    ap.add_argument("--outdir", default=None,
                    help="output directory for --sheet-from-prompts")
    ap.add_argument("--ask", action="store_true",
                    help="ask-first mode: clarify/refuse instead of drawing")
    ap.add_argument("--bundle", default=str(ROOT / "weights"))
    ap.add_argument("--no-quality", action="store_true",
                    help="disable the grounding auto-retry gate")
    args = ap.parse_args()

    if args.sheet_from_prompts:
        outdir = Path(args.outdir or "pxg_batch_out")
        outdir.mkdir(parents=True, exist_ok=True)
        lines = [l.strip() for l in Path(args.sheet_from_prompts)
                 .read_text().splitlines() if l.strip()]
        pipe = PXGPipeline(args.bundle)
        grids = []
        results = []
        for i, p in enumerate(lines):
            g, m = pipe.generate_pixels(p, seed=args.seed,
                                        enforce_quality=not args.no_quality,
                                        temperature=args.temperature,
                                        top_k=args.top_k)
            results.append({"prompt": p, "gate": m.get("gate")})
            if g is None:
                print(json.dumps({"i": i, "prompt": p, **m}, indent=2))
                continue
            from pxg_tiny.render import save_png
            f = outdir / f"{i+1:02d}_{p.replace(' ', '_')[:24]}.png"
            save_png(g, f, scale=8)
            grids.append(g)
            results[-1]["file"] = str(f)
        if grids:
            sheet(grids, outdir / "sheet.png")
        print(json.dumps({"n_prompts": len(lines), "n_rendered": len(grids),
                          "outdir": str(outdir), "results": results}, indent=2))
        return

    if not args.prompt:
        ap.error("give me a prompt, e.g. pxg-tiny \"a gold sword\"")

    pipe = PXGPipeline(args.bundle)

    if args.ask:
        label, msg = pipe.should_ask(args.prompt)
        print(json.dumps({"label": label, "message": msg}, indent=2))
        return

    out = Path(args.out) if args.out else Path(
        args.prompt.replace(" ", "_")[:24].replace("/", "_") + ".png")
    out.parent.mkdir(parents=True, exist_ok=True)

    grids, metas = [], []
    if args.variations > 1:
        res = pipe.variations(args.prompt, k=args.variations,
                              start_seed=args.seed,
                              enforce_quality=not args.no_quality,
                              temperature=args.temperature, top_k=args.top_k)
        grids = [g for g, _ in res if g is not None]
        metas = [m for _, m in res]
    else:
        g, m = pipe.generate_pixels(args.prompt, seed=args.seed,
                                    temperature=args.temperature,
                                    top_k=args.top_k,
                                    enforce_quality=not args.no_quality)
        if g is None:
            print(json.dumps(m, indent=2))
            return
        grids, metas = [g], [m]

    sheet(grids, out)
    print(json.dumps({"prompt": args.prompt, "out": str(out),
                      "n": len(grids), "meta": metas[0]}, indent=2))


if __name__ == "__main__":
    main()