David310/Detect_AI-generated_Image
4
1import functools2import torch3import torch.nn as nn4from networks.base_model import BaseModel, init_weights5import sys6from models import get_model7 8class Trainer(BaseModel):9 def name(self):10 return 'Trainer'11 12 def __init__(self, opt):13 super(Trainer, self).__init__(opt)14 self.opt = opt 15 self.model = get_model(opt.arch)16 torch.nn.init.normal_(self.model.fc.weight.data, 0.0, opt.init_gain)17 18 if opt.fix_backbone:19 params = []20 for name, p in self.model.named_parameters():21 if name=="fc.weight" or name=="fc.bias": 22 params.append(p) 23 else:24 p.requires_grad = False25 else:26 print("Your backbone is not fixed. Are you sure you want to proceed? If this is a mistake, enable the --fix_backbone command during training and rerun")27 import time 28 time.sleep(3)29 params = self.model.parameters()30 31 32 33 if opt.optim == 'adam':34 self.optimizer = torch.optim.AdamW(params, lr=opt.lr, betas=(opt.beta1, 0.999), weight_decay=opt.weight_decay)35 elif opt.optim == 'sgd':36 self.optimizer = torch.optim.SGD(params, lr=opt.lr, momentum=0.0, weight_decay=opt.weight_decay)37 else:38 raise ValueError("optim should be [adam, sgd]")39 40 self.loss_fn = nn.BCEWithLogitsLoss()41 42 self.model.to(opt.gpu_ids[0])43 44 45 def adjust_learning_rate(self, min_lr=1e-6):46 for param_group in self.optimizer.param_groups:47 param_group['lr'] /= 10.48 if param_group['lr'] < min_lr:49 return False50 return True51 52 53 def set_input(self, input):54 self.input = input[0].to(self.device)55 self.label = input[1].to(self.device).float()56 57 58 def forward(self):59 self.output = self.model(self.input)60 self.output = self.output.view(-1).unsqueeze(1)61 62 63 def get_loss(self):64 return self.loss_fn(self.output.squeeze(1), self.label)65 66 def optimize_parameters(self):67 self.forward()68 self.loss = self.loss_fn(self.output.squeeze(1), self.label) 69 self.optimizer.zero_grad()70 self.loss.backward()71 self.optimizer.step()72 73 74 75 