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

Function main

diffportrait360_release/code/train.py:169–486  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

167
168
169def main(args):
170
171 # ******************************
172 # initialize training
173 # ******************************
174 # assign rank os.environ['MASTER_ADDR'] = 'localhost'
175 os.environ['MASTER_PORT'] = '12355'
176
177
178 args.world_size = int(os.environ['WORLD_SIZE'])
179
180 args.local_rank = int(os.environ['LOCAL_RANK'])
181 print("local_rank", args.local_rank)
182 #import pdb;pdb.set_trace()
183 args.rank = int(os.environ['RANK'])
184 args.device = torch.device("cuda", args.local_rank)
185 args.num_gpu = torch.cuda.device_count()
186 args.use_gpu = torch.cuda.is_available() and args.num_gpu > 0
187 #seg_model = load_model(args, args.model_path, True, False)
188 os.makedirs(args.local_image_dir,exist_ok=True)
189 os.makedirs(args.local_log_dir,exist_ok=True)
190 if args.rank == 0:
191 print(args)
192
193 # initial distribution comminucation
194 dist.init_process_group("nccl", rank=args.rank, world_size=args.world_size)
195 torch.backends.cuda.matmul.allow_tf32 = False # it doenst work once not trained in A100 64G gpu
196 torch.backends.cudnn.benchmark = True
197
198 # set seed for reproducibility
199 set_seed(args.seed)
200
201 # visdom / tensorboard
202 if args.rank == 0:
203 tb_writer = SummaryWriter(log_dir=args.local_log_dir)
204 else:
205 tb_writer = None
206
207 # ******************************
208 # create model
209 # ******************************
210 #import pdb;pdb.set_trace()
211 model = create_model(args.model_config).cpu()
212
213 model.sd_locked = args.sd_locked
214 model.only_mid_control = args.only_mid_control
215 model.to(args.local_rank)
216 #seg_model.to(args.local_rank)
217 if args.local_rank == 0:
218 print('Total base parameters {:.02f}M'.format(count_param([model])))
219 model_ema = None
220
221 # ******************************
222
223 # ******************************
224 # load pre-trained models
225 # ******************************
226 optimizer_state_dict = None

Callers 1

train.pyFile · 0.70

Calls 15

set_seedFunction · 0.90
create_modelFunction · 0.90
count_paramFunction · 0.90
print_peak_memoryFunction · 0.90
merge_lists_by_indexFunction · 0.90
save_checkpoint_emaFunction · 0.90
stepMethod · 0.80
appendMethod · 0.80
meanMethod · 0.80
load_state_dictFunction · 0.70
get_cond_controlFunction · 0.70

Tested by

no test coverage detected