(gpu, opt, output_dir, noises_init)
| 544 | |
| 545 | |
| 546 | def 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) |
no test coverage detected