MCPcopy Create free account
hub / github.com/VisionRush/DeepFakeDefenders / create_env

Method create_env

core/mengine.py:44–68  ·  view source on GitHub ↗
(self, cfg)

Source from the content-addressed store, hash-verified

42 self.SyncBN = SyncBatchNorm
43
44 def create_env(self, cfg):
45 # create network
46 self.netloc_ = load_model(cfg.network.name, cfg.network.class_num, self.SyncBN)
47 print(self.netloc_)
48
49 self.netloc_.cuda()
50 if self.DDP:
51 if self.SyncBN:
52 self.netloc_ = torch.nn.SyncBatchNorm.convert_sync_batchnorm(self.netloc_)
53 self.netloc_ = DDP(self.netloc_,
54 device_ids=[self.local_rank],
55 broadcast_buffers=True,
56 )
57
58 # create loss function
59 self.criterion_ = nn.CrossEntropyLoss().cuda()
60
61 # create optimizer
62 self.optimizer_ = torch.optim.AdamW(self.netloc_.parameters(), lr=cfg.optimizer.lr,
63 betas=(cfg.optimizer.beta1, cfg.optimizer.beta2), eps=cfg.optimizer.eps,
64 weight_decay=cfg.optimizer.weight_decay)
65
66 # create scheduler
67 self.scheduler_ = torch.optim.lr_scheduler.CosineAnnealingLR(self.optimizer_, cfg.train.epoch_num,
68 eta_min=cfg.scheduler.min_lr)
69
70 def train_multi_class(self, train_loader, epoch_idx, ema_start):
71 starttime = datetime.datetime.now()

Callers 2

main_train.pyFile · 0.80

Calls 1

load_modelFunction · 0.90

Tested by

no test coverage detected