"""Local or Hub inference for Image21-INT4.""" import argparse from pathlib import Path import torch from scripts.editing_protocol import DEFAULT_SEEDS, edit_dimensions, load_input, validate_seeds from scripts.runtime import load_pipeline def main(): parser = argparse.ArgumentParser() parser.add_argument('--model', default='models/int4') parser.add_argument('--prompt', required=True) parser.add_argument('--input', nargs='*', default=[]) parser.add_argument('--output', default='output.png') parser.add_argument('--seed', type=int) parser.add_argument('--source-seed', type=int) parser.add_argument('--width', type=int) parser.add_argument('--height', type=int) parser.add_argument('--resolution', type=int, default=1024) parser.add_argument('--steps', type=int, default=40) parser.add_argument('--offload', choices=['auto', 'model', 'group', 'resident'], default='auto') parser.add_argument('--device', default=None) args = parser.parse_args() if (args.width is None) != (args.height is None): parser.error('Specify both --width and --height, or neither') if args.steps <= 0 or any(value <= 0 or value % 32 for value in (args.width, args.height, args.resolution) if value is not None): parser.error('Positive steps and dimensions divisible by 32 are required') if Path(args.output).suffix.lower() != '.png': parser.error('Use a .png output so RGBA channels are retained') 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) case = {'source_seed': args.source_seed} if images else {} try: validate_seeds([case], [seed]) except ValueError as exc: parser.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)) pipe = load_pipeline(args.model, offload=args.offload, device=args.device, local_files_only=Path(args.model).exists()) 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) policy = pipe.image21_runtime print(f'{args.output} (seed={seed}, size={image.size}, mode={image.mode}, ' f'device={policy["device"]}, offload={policy["offload"]})') if __name__ == '__main__': main()