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