File size: 2,842 Bytes
9116984
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
"""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()