-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathinference_sdxl.py
More file actions
executable file
·134 lines (125 loc) · 3.16 KB
/
Copy pathinference_sdxl.py
File metadata and controls
executable file
·134 lines (125 loc) · 3.16 KB
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
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
import argparse
import yaml
from nb_utils.eval_sets import base_set, live_set, object_set, merge_test_set, merge_base_set
from nb_utils.configs import live_object_data
from moft.inferencer_sdxl import inferencers
import warnings
warnings.filterwarnings('ignore')
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument(
"--inference_type",
type=str,
required=True,
)
parser.add_argument(
"--config_path",
type=str,
required=True,
help="Path to hparams.yml"
)
parser.add_argument(
"--checkpoint_idx",
type=str,
default=None,
required=False,
)
parser.add_argument(
"--t",
type=float,
default=None,
required=False,
help="t value for merging"
)
parser.add_argument(
"--parameter",
type=float,
default=None,
required=False,
help="parameter value for postprocessing"
)
parser.add_argument(
"--postprocessing_method",
type=str,
default="curve_over_id",
required=False,
help="method for postprocessing the merged matrices"
)
parser.add_argument(
"--num_images_per_medium_prompt",
type=int,
default=1,
help="Number of generated images for each medium prompt",
)
parser.add_argument(
"--num_images_per_base_prompt",
type=int,
default=10,
help="Number of generated images for each base prompt",
)
parser.add_argument(
"--batch_size_medium",
type=int,
default=1,
)
parser.add_argument(
"--batch_size_base",
type=int,
default=10,
)
parser.add_argument(
"--num_inference_steps",
type=int,
default=50
)
parser.add_argument(
"--guidance_scale",
type=float,
default=7.0
)
parser.add_argument(
"--replace_inference_output",
action='store_true',
default=False
)
parser.add_argument(
"--version",
type=int,
default=0
)
parser.add_argument(
"--seed",
type=int,
default=0
)
parser.add_argument(
"--moft_layers_concept_path",
type=str,
default=None,
required=False,
help="Path to moft layers concept"
)
parser.add_argument(
"--moft_layers_style_path",
type=str,
default=None,
required=False,
help="Path to moft layers style"
)
return parser.parse_args()
def main(args):
with open(args.config_path, 'r', encoding='utf-8') as config_file:
config = yaml.safe_load(config_file)
if live_object_data[config['class_name']] == 'live':
evaluation_set = live_set
else:
evaluation_set = object_set
if args.t is not None or args.inference_type == 'moft_direct_merge':
evaluation_set = merge_test_set
print(evaluation_set)
inferencer = inferencers[args.inference_type](config, args, evaluation_set, merge_base_set)
inferencer.setup()
inferencer.generate()
if __name__ == '__main__':
args = parse_args()
main(args)