PixelModel-v2 / main.py
wop's picture
Upload 17 files
e465a2f verified
Raw
History Blame Contribute Delete
1.36 kB
"""
main.py - PixelModel v1 inference from model.png (the canonical model).
Usage:
python main.py "a red double decker bus"
python main.py "a cat on a couch" --out cat.png --res 64 --scale 4
"""
import argparse
import numpy as np
import torch
from PIL import Image
from model import NATIVE_RES, forward, load_model
def main():
p = argparse.ArgumentParser()
p.add_argument("prompt")
p.add_argument("--model", default="model.png")
p.add_argument("--out", default="out.png")
p.add_argument("--res", type=int, default=NATIVE_RES,
help="generation resolution (decoder is resolution-free)")
p.add_argument("--scale", type=int, default=4, help="nearest-neighbor upscale for viewing")
args = p.parse_args()
weights = load_model(args.model)
with torch.no_grad():
result = forward(weights, args.prompt, res=args.res)[0]
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}")
print(f"output : {args.out} ({args.res}x{args.res} native, x{args.scale} view)")
if __name__ == "__main__":
main()