Image21-INT4 / scripts /infer.py
ixim's picture
Release verified Image21-INT4 conversion
9116984 verified
Raw History Blame Contribute Delete
2.84 kB
"""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()