CoolFace
Modelpublic

CLYang617/RemoteSensingChangeDetection-RSCD.HA2F

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
utils.py82 linesDownload Raw Back to model
1import torch2import torch.nn.functional as F3import numpy as np4import torch.nn as nn5import random6 7 8def weight_init(module):9    for n, m in module.named_children():10        print('initialize: '+n)11        if isinstance(m, nn.Conv2d):12            nn.init.kaiming_normal_(m.weight, mode='fan_in', nonlinearity='relu')13            if m.bias is not None:14                nn.init.zeros_(m.bias)15        elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):16            nn.init.ones_(m.weight)17            if m.bias is not None:18                nn.init.zeros_(m.bias)19        elif isinstance(m, nn.Linear):20            nn.init.kaiming_normal_(m.weight, mode='fan_in', nonlinearity='relu')21            if m.bias is not None:22                nn.init.zeros_(m.bias)23        elif isinstance(m, nn.Sequential):24            for f, g in m.named_children():25                print('initialize: ' + f)26                if isinstance(g, nn.Conv2d):27                    nn.init.kaiming_normal_(g.weight, mode='fan_in', nonlinearity='relu')28                    if g.bias is not None:29                        nn.init.zeros_(g.bias)30                elif isinstance(g, (nn.BatchNorm2d, nn.GroupNorm)):31                    nn.init.ones_(g.weight)32                    if g.bias is not None:33                        nn.init.zeros_(g.bias)34                elif isinstance(g, nn.Linear):35                    nn.init.kaiming_normal_(g.weight, mode='fan_in', nonlinearity='relu')36                    if g.bias is not None:37                        nn.init.zeros_(g.bias)38        elif isinstance(m, nn.AdaptiveAvgPool2d) or isinstance(m, nn.AdaptiveMaxPool2d) or isinstance(m, nn.ModuleList) or isinstance(m, nn.BCELoss):39            a=140        else:41            pass42 43 44def init_seed(seed):45    torch.manual_seed(seed)46    torch.cuda.manual_seed(seed)47    random.seed(seed)48    np.random.seed(seed)49 50 51def BCEDiceLoss(inputs, targets):52    # print(inputs.shape, targets.shape)53    bce = F.binary_cross_entropy(inputs, targets)54    inter = (inputs * targets).sum()55    eps = 1e-556    dice = (2 * inter + eps) / (inputs.sum() + targets.sum() + eps)57    # print(bce.item(), inter.item(), inputs.sum().item(), dice.item())58    return bce + 1 - dice59 60 61def BCE(inputs, targets):62    # print(inputs.shape, targets.shape)63    bce = F.binary_cross_entropy(inputs, targets)64    return bce65 66 67def adjust_learning_rate(args, optimizer, epoch, iter, max_batches, lr_factor=1):68    if args.lr_mode == 'step':69        lr = args.lr * (0.1 ** (epoch // args.step_loss))70    elif args.lr_mode == 'poly':71        cur_iter = iter72        max_iter = max_batches * args.max_epochs73        lr = args.lr * (1 - cur_iter * 1.0 / max_iter) ** 0.974    else:75        raise ValueError('Unknown lr mode {}'.format(args.lr_mode))76    if epoch == 0 and iter < 200:77        lr = args.lr * 0.9 * (iter + 1) / 200 + 0.1 * args.lr  # warm_up78    lr *= lr_factor79    for param_group in optimizer.param_groups:80        param_group['lr'] = lr81    return lr82