PixelModel-v2 / INFERENCE.py
wop's picture
Upload 17 files
e465a2f verified
Raw
History Blame
3.56 kB
"""
INFERENCE.py - run PixelModel v2 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 v2 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()