File size: 1,901 Bytes
63a1291
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""PixelModel v3 inference from the canonical model.png.

    python main.py "a red double decker bus" --out bus.png
    python main.py "a beach with palm trees" --res 256

model.png is the model. This script decodes it, tokenises the prompt with the
shipped vocab.json, and paints an image at any resolution. Fully deterministic,
so it matches INFERENCE.py (which loads model.safetensors) bit for bit.
"""

from __future__ import annotations

import argparse

import numpy as np
import torch
from PIL import Image

from model import (
    load_config, load_model_png, load_vocab, encode_caption, make_coord_grid,
)


def render(model, cfg, vocab, prompt, res, device):
    tokens = encode_caption(prompt, vocab, cfg.max_tokens)
    tokens = torch.from_numpy(tokens).long().unsqueeze(0).to(device)
    coords = make_coord_grid(res, res, device=device, dtype=torch.float32).unsqueeze(0)
    with torch.no_grad():
        rgb = model(tokens, coords)
    img = (rgb.clamp(0, 1).reshape(res, res, 3).cpu().numpy() * 255.0).round().astype(np.uint8)
    return img


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("prompt")
    ap.add_argument("--out", default="out.png")
    ap.add_argument("--res", type=int, default=128, help="output resolution (native 128)")
    ap.add_argument("--png", default="model.png", help="the model")
    ap.add_argument("--config", default="config.json")
    ap.add_argument("--vocab", default="vocab.json")
    ap.add_argument("--device", default="cpu")
    args = ap.parse_args()

    cfg = load_config(args.config)
    model = load_model_png(args.png, cfg, map_location=args.device)
    vocab = load_vocab(args.vocab)

    img = render(model, cfg, vocab, args.prompt, args.res, args.device)
    Image.fromarray(img, "RGB").save(args.out)
    print(f'[main] "{args.prompt}" @ {args.res}x{args.res} -> {args.out}')


if __name__ == "__main__":
    main()