MCPcopy Create free account
hub / github.com/FreedomGu/Diffportrait360 / main

Function main

diffportrait360_release/code/inference.py:138–213  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

136 for idm, tensor in enumerate(gene_img_list):
137 writer_gen.append_data(generated_imgs[idm])
138def main(args):
139 # ******************************
140 # initializing
141 # ******************************
142 args.device = torch.device("cuda")
143 args.num_gpu = torch.cuda.device_count()
144 args.use_gpu = torch.cuda.is_available() and args.num_gpu > 0
145 #seg_model = load_model(args, args.model_path, True, False)
146 os.makedirs(args.local_image_dir,exist_ok=True)
147 print(args)
148 set_seed(args.seed)
149 # ******************************
150 # create model
151 # ******************************
152 model = create_model(args.model_config).cpu()
153 model.sd_locked = args.sd_locked
154 model.only_mid_control = args.only_mid_control
155 model.to(args.device)
156 print('Total base parameters {:.02f}M'.format(count_param([model])))
157 # ******************************
158 # load pre-trained models
159 # ******************************
160 ckpt_path = args.resume_dir
161 print('loading state dict from {} ...'.format(ckpt_path))
162 load_state_dict(model, ckpt_path, strict=True)
163 torch.cuda.empty_cache()
164 # ******************************
165 # create dataset and dataloader
166 # ******************************
167 image_transform = T.Compose([
168 T.ToTensor(),
169 T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
170 ])
171 if args.test_dataset == 'back_head_generation':
172 test_dataset_cls = getattr(full_head_clean, args.test_dataset)
173 test_image_dataset = test_dataset_cls(
174 image_transform = image_transform,
175 inference_image_dataset = args.inference_image_path,
176 condition_path = args.condition_path
177 )
178 elif args.test_dataset == "full_head_clean_inference_final_face":
179 test_dataset_cls = getattr(full_head_clean, args.test_dataset)
180 test_image_dataset = test_dataset_cls(
181 image_transform=image_transform,
182 condition_path = args.condition_path,
183 inference_image_dataset = args.inference_image_path,
184 initial_image_path = args.initial_image_path,
185 #extra_appearance_num = args.extra_appearance_num,
186 #mask_condition = args.mask_condition ,
187 )
188 else:
189 print("find the appropriate dataset class!")
190 return
191 test_image_dataloader = DataLoader(test_image_dataset,
192 batch_size=1,
193 num_workers=0,
194 #pin_memory=True,
195 shuffle=False)

Callers 1

inference.pyFile · 0.70

Calls 6

set_seedFunction · 0.90
create_modelFunction · 0.90
count_paramFunction · 0.90
print_peak_memoryFunction · 0.90
load_state_dictFunction · 0.70
visualizeFunction · 0.70

Tested by

no test coverage detected