Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
trainer.py75 linesDownload Raw Back to networks
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