File size: 3,453 Bytes
c500926
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
INFERENCE.py - run PixelModel v1 from model.safetensors, self-contained.

Only needs torch + safetensors + pillow + numpy; no other file from this
repo. Produces byte-identical output to main.py (which loads model.png)
for the same prompt and resolution.

Usage:
  python INFERENCE.py "a red double decker bus"
  python INFERENCE.py "a cat on a couch" --model model.safetensors --out cat.png --res 64 --scale 4
"""

import argparse
import os
import sys

import numpy as np
import torch
from PIL import Image
from safetensors.torch import load_file

EMB_DIM    = 64
NATIVE_RES = 64
FREQS      = (1.0, 2.0, 4.0, 8.0)


def _fnv1a(data: bytes) -> int:
    h = 0x811C9DC5
    for byte in data:
        h ^= byte
        h = (h * 0x01000193) & 0xFFFFFFFF
    return h


def prompt_to_embedding(prompt: str) -> torch.Tensor:
    text = "".join(c if c.isalnum() or c == " " else " " for c in prompt.lower())
    text = " ".join(text.split())
    vec = np.zeros(EMB_DIM, dtype=np.float32)
    padded = f" {text} "
    for i in range(len(padded) - 2):
        h = _fnv1a(padded[i:i + 3].encode("utf-8"))
        vec[h % EMB_DIM] += 1.0 if (h >> 16) & 1 else -1.0
    for word in text.split():
        h = _fnv1a(b"w:" + word.encode("utf-8"))
        vec[h % EMB_DIM] += 2.0 if (h >> 16) & 1 else -2.0
    norm = np.linalg.norm(vec)
    if norm > 0:
        vec /= norm
    return torch.from_numpy(vec)


def coord_features(res: int) -> torch.Tensor:
    axis = torch.linspace(-1.0, 1.0, res)
    yy, xx = torch.meshgrid(axis, axis, indexing="ij")
    x, y = xx.reshape(-1), yy.reshape(-1)
    feats = [x, y]
    for f in FREQS:
        feats += [torch.sin(f * torch.pi * x), torch.cos(f * torch.pi * x),
                  torch.sin(f * torch.pi * y), torch.cos(f * torch.pi * y)]
    return torch.stack(feats, dim=1)


def forward(w: dict, prompt: str, res: int) -> torch.Tensor:
    emb = prompt_to_embedding(prompt).unsqueeze(0)
    z = torch.tanh(emb @ w["T1"].T + w["b1"])
    z = torch.tanh(z @ w["T2"].T + w["b2"])
    feats = coord_features(res)
    P = feats.shape[0]
    inp = torch.cat([z.expand(P, -1), feats], dim=1)
    h = torch.tanh(inp @ w["D1"].T + w["bd1"])
    h = torch.tanh(h @ w["D2"].T + w["bd2"])
    rgb = torch.sigmoid(h @ w["D3"].T + w["bd3"])
    return rgb.reshape(res, res, 3)


def main():
    p = argparse.ArgumentParser(description="PixelModel v1 inference (safetensors)")
    p.add_argument("prompt")
    p.add_argument("--model", default="model.safetensors")
    p.add_argument("--out",   default="out.png")
    p.add_argument("--res",   type=int, default=NATIVE_RES)
    p.add_argument("--scale", type=int, default=4)
    args = p.parse_args()

    if not os.path.exists(args.model):
        sys.exit(f"Model not found: {args.model}\n"
                 f"Run: python convert_to_safetensors.py   to create it from model.png.")

    weights = load_file(args.model)
    with torch.no_grad():
        result = forward(weights, args.prompt, args.res)

    arr = (result.numpy() * 255).clip(0, 255).astype(np.uint8)
    img = Image.fromarray(arr, mode="RGB")
    if args.scale > 1:
        img = img.resize((args.res * args.scale,) * 2, Image.NEAREST)
    img.save(args.out)
    print(f"prompt : '{args.prompt}'")
    print(f"model  : {args.model}  ({os.path.getsize(args.model)} bytes, safetensors)")
    print(f"output : {args.out}  ({args.res}x{args.res} native, x{args.scale} view)")


if __name__ == "__main__":
    main()