PixelModel-v3 / INFERENCE.py
TobiasLogic's picture
PixelModel v3: SIREN+FiLM CPPN, 919K params, beats v1 FID
63a1291 verified
Raw
History Blame
2.12 kB
"""PixelModel v3 inference from model.safetensors (standalone).
Needs only: model.safetensors, model.py, vocab.json, torch. No model.png,
no config.json - the architecture config is read from the safetensors header.
Produces output identical to main.py for the same prompt and resolution.
python convert_to_safetensors.py
python INFERENCE.py "a red double decker bus" --out bus.png
"""
from __future__ import annotations
import argparse
import json
import numpy as np
import torch
from PIL import Image
from safetensors import safe_open
from model import (
ModelConfig, PixelModelV3, load_vocab, encode_caption, make_coord_grid,
)
def load_from_safetensors(path, device):
with safe_open(path, framework="pt", device="cpu") as f:
meta = f.metadata() or {}
cfg = ModelConfig.from_json(json.loads(meta["config"]))
model = PixelModelV3(cfg)
state = {k: f.get_tensor(k) for k in f.keys()}
model.load_state_dict(state)
model.to(device).eval()
return model, cfg
def main():
ap = argparse.ArgumentParser()
ap.add_argument("prompt")
ap.add_argument("--out", default="out.png")
ap.add_argument("--res", type=int, default=128)
ap.add_argument("--safetensors", default="model.safetensors")
ap.add_argument("--vocab", default="vocab.json")
ap.add_argument("--device", default="cpu")
args = ap.parse_args()
model, cfg = load_from_safetensors(args.safetensors, args.device)
vocab = load_vocab(args.vocab)
tokens = encode_caption(args.prompt, vocab, cfg.max_tokens)
tokens = torch.from_numpy(tokens).long().unsqueeze(0).to(args.device)
coords = make_coord_grid(args.res, args.res, device=args.device,
dtype=torch.float32).unsqueeze(0)
with torch.no_grad():
rgb = model(tokens, coords)
img = (rgb.clamp(0, 1).reshape(args.res, args.res, 3).cpu().numpy() * 255.0
).round().astype(np.uint8)
Image.fromarray(img, "RGB").save(args.out)
print(f'[inference] "{args.prompt}" @ {args.res}x{args.res} -> {args.out}')
if __name__ == "__main__":
main()