| """ | |
| 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() | |