File size: 2,500 Bytes
b451f49
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
INFERENCE.py — run PixelModel from its model.safetensors weights.

This is the standalone entry point for anyone who just wants to load the
safetensors weights and generate an image, without needing model.png or
the rest of this repo's training code.

Usage:
  python INFERENCE.py "a red circle"
  python INFERENCE.py "a red circle" --model model.safetensors --out out.png --scale 8
"""

import argparse
import os
import sys

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

PROMPT_DIM = 32
OUT_SIZE = 32


def prompt_to_embedding(prompt: str) -> torch.Tensor:
    """Deterministic char-level embedding -> PROMPT_DIM vector."""
    vec = torch.zeros(PROMPT_DIM)
    for i, ch in enumerate(prompt.lower()):
        idx = i % PROMPT_DIM
        vec[idx] += ord(ch) / 127.0
    norm = vec.norm()
    if norm > 0:
        vec = vec / norm
    return vec


def forward(weights: dict, prompt: str) -> torch.Tensor:
    emb = prompt_to_embedding(prompt)
    x = torch.tanh(weights["W1"] @ emb)
    x = torch.tanh(weights["W2"] @ x)
    x = torch.sigmoid(weights["W3"] @ x)
    return x.reshape(OUT_SIZE, OUT_SIZE, 3)


def generate(prompt: str, model_path: str, out_path: str, scale: int = 8):
    if not os.path.exists(model_path):
        sys.exit(
            f"Model not found: {model_path}\n"
            f"Run: python convert_to_safetensors.py   to create one from model.png first."
        )

    weights = load_file(model_path)

    with torch.no_grad():
        result = forward(weights, prompt)

    arr = (result.numpy() * 255).clip(0, 255).astype(np.uint8)
    img = Image.fromarray(arr, mode="RGB")

    if scale > 1:
        img = img.resize((OUT_SIZE * scale, OUT_SIZE * scale), Image.NEAREST)

    img.save(out_path)
    print(f"prompt : '{prompt}'")
    print(f"model  : {model_path}  ({os.path.getsize(model_path)} bytes, safetensors)")
    print(f"output : {out_path}  ({OUT_SIZE * scale}x{OUT_SIZE * scale} px)")


if __name__ == "__main__":
    p = argparse.ArgumentParser(description="PixelModel inference (safetensors)")
    p.add_argument("prompt", help="Text prompt")
    p.add_argument("--model", default="model.safetensors", help="Path to safetensors weights")
    p.add_argument("--out", default="out.png", help="Output image path")
    p.add_argument("--scale", type=int, default=8, help="Upscale factor for output (default 8 -> 256x256)")
    args = p.parse_args()
    generate(args.prompt, args.model, args.out, args.scale)