MCPcopy Create free account
hub / github.com/Zhiyuan-R/Tiger-Diffusion / train

Function train

train_generation.py:546–751  ·  view source on GitHub ↗
(gpu, opt, output_dir, noises_init)

Source from the content-addressed store, hash-verified

544
545
546def train(gpu, opt, output_dir, noises_init):
547
548 set_seed(opt)
549 logger = setup_logging(output_dir)
550 if opt.distribution_type == 'multi':
551 should_diag = gpu==0
552 else:
553 should_diag = True
554 if should_diag:
555 outf_syn, = setup_output_subdirs(output_dir, 'syn')
556
557 if opt.distribution_type == 'multi':
558 if opt.dist_url == "env://" and opt.rank == -1:
559 opt.rank = int(os.environ["RANK"])
560
561 base_rank = opt.rank * opt.ngpus_per_node
562 opt.rank = base_rank + gpu
563 dist.init_process_group(backend=opt.dist_backend, init_method=opt.dist_url,
564 world_size=opt.world_size, rank=opt.rank)
565
566 opt.bs = int(opt.bs / opt.ngpus_per_node)
567 opt.workers = 0
568
569 opt.saveIter = int(opt.saveIter / opt.ngpus_per_node)
570 opt.diagIter = int(opt.diagIter / opt.ngpus_per_node)
571 opt.vizIter = int(opt.vizIter / opt.ngpus_per_node)
572
573
574 ''' data '''
575 train_dataset, _ = get_dataset(opt.dataroot, opt.npoints, opt.category)
576 dataloader, _, train_sampler, _ = get_dataloader(opt, train_dataset, None)
577
578
579 '''
580 create networks
581 '''
582
583 betas = get_betas(opt.schedule_type, opt.beta_start, opt.beta_end, opt.time_num)
584 model = Model(opt, betas, opt.loss_type, opt.model_mean_type, opt.model_var_type)
585
586 if opt.distribution_type == 'multi': # Multiple processes, single GPU per process
587 def _transform_(m):
588 return nn.parallel.DistributedDataParallel(
589 m, device_ids=[gpu], output_device=gpu)
590
591 torch.cuda.set_device(gpu)
592 model.cuda(gpu)
593 model.multi_gpu_wrapper(_transform_)
594
595
596 elif opt.distribution_type == 'single':
597 def _transform_(m):
598 return nn.parallel.DataParallel(m)
599 model = model.cuda()
600 model.multi_gpu_wrapper(_transform_)
601
602 elif gpu is not None:
603 torch.cuda.set_device(gpu)

Callers 1

mainFunction · 0.85

Calls 15

multi_gpu_wrapperMethod · 0.95
get_loss_iterMethod · 0.95
all_klMethod · 0.95
evalMethod · 0.95
gen_samplesMethod · 0.95
gen_sample_trajMethod · 0.95
trainMethod · 0.95
set_seedFunction · 0.85
setup_loggingFunction · 0.85
setup_output_subdirsFunction · 0.85
get_dataloaderFunction · 0.85
getGradNormFunction · 0.85

Tested by

no test coverage detected