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