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