Image21-INT8 / scripts /infer.py
ixim's picture
Release verified Image21-INT8 conversion (part 2)
e9190ff verified
Raw History Blame Contribute Delete
3.01 kB
"""Local or Hub inference for the published Diffusers checkpoint."""
import argparse
from pathlib import Path
import torch
from scripts.runtime import load_int8_pipeline
from scripts.editing_protocol import load_input, validate_seeds, DEFAULT_SEEDS, edit_dimensions
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--model', default='models/int8')
ap.add_argument('--prompt', required=True)
ap.add_argument('--input', nargs='*', default=[])
ap.add_argument('--output', default='output.png')
ap.add_argument('--seed', type=int, help='Default: 42 for generation, 1000042 for editing')
ap.add_argument('--source-seed', type=int, help='Known source generation seed; reject editing noise replay')
ap.add_argument('--width', type=int)
ap.add_argument('--height', type=int)
ap.add_argument('--resolution', type=int, default=1024,
help='Editing reference/target area scale, independent of explicit output width')
ap.add_argument('--steps', type=int, default=40)
ap.add_argument('--vae-tiling', action='store_true',
help='Save VAE memory; may introduce the stripe artifacts documented in EDITING_UPDATE.md')
args = ap.parse_args()
if ((args.width is None) != (args.height is None)):
ap.error('Specify both --width and --height, or neither')
if args.steps <= 0 or any(v <= 0 or v % 32 for v in
(args.width, args.height, args.resolution) if v is not None):
ap.error('Positive steps and dimensions divisible by 32 are required')
if Path(args.output).suffix.lower() != '.png':
ap.error('Use a .png output to retain the RGBA channel')
images = [load_input(path) for path in args.input]
seed = args.seed if args.seed is not None else (DEFAULT_SEEDS[0] if images else 42)
try:
validate_seeds([{'source_seed': args.source_seed}] if images else [], [seed])
except ValueError as exc:
ap.error(str(exc))
if args.width is None:
args.width, args.height = (edit_dimensions(images[-1].size, args.resolution)
if images else (args.resolution, args.resolution))
if Path(args.model).exists():
from scripts.benchmark import load_pipeline
pipe = load_pipeline(args.model)
else:
pipe = load_int8_pipeline(args.model)
if args.vae_tiling:
pipe.vae.enable_tiling()
with torch.inference_mode():
image = pipe(prompt=args.prompt, image=images or None,
width=args.width, height=args.height, output_resolution=args.resolution,
num_inference_steps=args.steps, true_cfg_scale=1.0, use_kv_cache=True,
generator=torch.Generator('cpu').manual_seed(seed)).images[0]
Path(args.output).parent.mkdir(parents=True, exist_ok=True)
image.save(args.output)
print(f'{args.output} (seed={seed}, actual_size={image.size}, mode={image.mode})')
if __name__ == '__main__':
main()