init
This commit is contained in:
@@ -0,0 +1,97 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
def update_learning_rate(P, optimizer, cur_epoch, n, n_total):
|
||||
|
||||
cur_epoch = cur_epoch - 1
|
||||
|
||||
lr = P.lr_init
|
||||
if P.optimizer == 'sgd' or 'lars':
|
||||
DECAY_RATIO = 0.1
|
||||
elif P.optimizer == 'adam':
|
||||
DECAY_RATIO = 0.3
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
if P.warmup > 0:
|
||||
cur_iter = cur_epoch * n_total + n
|
||||
if cur_iter <= P.warmup:
|
||||
lr *= cur_iter / float(P.warmup)
|
||||
|
||||
if cur_epoch >= 0.5 * P.epochs:
|
||||
lr *= DECAY_RATIO
|
||||
if cur_epoch >= 0.75 * P.epochs:
|
||||
lr *= DECAY_RATIO
|
||||
for param_group in optimizer.param_groups:
|
||||
param_group['lr'] = lr
|
||||
return lr
|
||||
|
||||
|
||||
def _cross_entropy(input, targets, reduction='mean'):
|
||||
targets_prob = F.softmax(targets, dim=1)
|
||||
xent = (-targets_prob * F.log_softmax(input, dim=1)).sum(1)
|
||||
if reduction == 'sum':
|
||||
return xent.sum()
|
||||
elif reduction == 'mean':
|
||||
return xent.mean()
|
||||
elif reduction == 'none':
|
||||
return xent
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
def _entropy(input, reduction='mean'):
|
||||
return _cross_entropy(input, input, reduction)
|
||||
|
||||
|
||||
def cross_entropy_soft(input, targets, reduction='mean'):
|
||||
targets_prob = F.softmax(targets, dim=1)
|
||||
xent = (-targets_prob * F.log_softmax(input, dim=1)).sum(1)
|
||||
if reduction == 'sum':
|
||||
return xent.sum()
|
||||
elif reduction == 'mean':
|
||||
return xent.mean()
|
||||
elif reduction == 'none':
|
||||
return xent
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
def kl_div(input, targets, reduction='batchmean'):
|
||||
return F.kl_div(F.log_softmax(input, dim=1), F.softmax(targets, dim=1),
|
||||
reduction=reduction)
|
||||
|
||||
|
||||
def target_nll_loss(inputs, targets, reduction='none'):
|
||||
inputs_t = -F.nll_loss(inputs, targets, reduction='none')
|
||||
logit_diff = inputs - inputs_t.view(-1, 1)
|
||||
logit_diff = logit_diff.scatter(1, targets.view(-1, 1), -1e8)
|
||||
diff_max = logit_diff.max(1)[0]
|
||||
|
||||
if reduction == 'sum':
|
||||
return diff_max.sum()
|
||||
elif reduction == 'mean':
|
||||
return diff_max.mean()
|
||||
elif reduction == 'none':
|
||||
return diff_max
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
def target_nll_c(inputs, targets, reduction='none'):
|
||||
conf = torch.softmax(inputs, dim=1)
|
||||
conf_t = -F.nll_loss(conf, targets, reduction='none')
|
||||
conf_diff = conf - conf_t.view(-1, 1)
|
||||
conf_diff = conf_diff.scatter(1, targets.view(-1, 1), -1)
|
||||
diff_max = conf_diff.max(1)[0]
|
||||
|
||||
if reduction == 'sum':
|
||||
return diff_max.sum()
|
||||
elif reduction == 'mean':
|
||||
return diff_max.mean()
|
||||
elif reduction == 'none':
|
||||
return diff_max
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,79 @@
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import diffdist.functional as distops
|
||||
|
||||
|
||||
def get_similarity_matrix(outputs, chunk=2, multi_gpu=False):
|
||||
'''
|
||||
Compute similarity matrix
|
||||
- outputs: (B', d) tensor for B' = B * chunk
|
||||
- sim_matrix: (B', B') tensor
|
||||
'''
|
||||
|
||||
if multi_gpu:
|
||||
outputs_gathered = []
|
||||
for out in outputs.chunk(chunk):
|
||||
gather_t = [torch.empty_like(out) for _ in range(dist.get_world_size())]
|
||||
gather_t = torch.cat(distops.all_gather(gather_t, out))
|
||||
outputs_gathered.append(gather_t)
|
||||
outputs = torch.cat(outputs_gathered)
|
||||
|
||||
sim_matrix = torch.mm(outputs, outputs.t()) # (B', d), (d, B') -> (B', B')
|
||||
|
||||
return sim_matrix
|
||||
|
||||
|
||||
def NT_xent(sim_matrix, temperature=0.5, chunk=2, eps=1e-8):
|
||||
'''
|
||||
Compute NT_xent loss
|
||||
- sim_matrix: (B', B') tensor for B' = B * chunk (first 2B are pos samples)
|
||||
'''
|
||||
|
||||
device = sim_matrix.device
|
||||
|
||||
B = sim_matrix.size(0) // chunk # B = B' / chunk
|
||||
|
||||
eye = torch.eye(B * chunk).to(device) # (B', B')
|
||||
sim_matrix = torch.exp(sim_matrix / temperature) * (1 - eye) # remove diagonal
|
||||
|
||||
denom = torch.sum(sim_matrix, dim=1, keepdim=True)
|
||||
sim_matrix = -torch.log(sim_matrix / (denom + eps) + eps) # loss matrix
|
||||
|
||||
loss = torch.sum(sim_matrix[:B, B:].diag() + sim_matrix[B:, :B].diag()) / (2 * B)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
def Supervised_NT_xent(sim_matrix, labels, temperature=0.5, chunk=2, eps=1e-8, multi_gpu=False):
|
||||
'''
|
||||
Compute NT_xent loss
|
||||
- sim_matrix: (B', B') tensor for B' = B * chunk (first 2B are pos samples)
|
||||
'''
|
||||
|
||||
device = sim_matrix.device
|
||||
|
||||
if multi_gpu:
|
||||
gather_t = [torch.empty_like(labels) for _ in range(dist.get_world_size())]
|
||||
labels = torch.cat(distops.all_gather(gather_t, labels))
|
||||
labels = labels.repeat(2)
|
||||
|
||||
logits_max, _ = torch.max(sim_matrix, dim=1, keepdim=True)
|
||||
sim_matrix = sim_matrix - logits_max.detach()
|
||||
|
||||
B = sim_matrix.size(0) // chunk # B = B' / chunk
|
||||
|
||||
eye = torch.eye(B * chunk).to(device) # (B', B')
|
||||
sim_matrix = torch.exp(sim_matrix / temperature) * (1 - eye) # remove diagonal
|
||||
|
||||
denom = torch.sum(sim_matrix, dim=1, keepdim=True)
|
||||
sim_matrix = -torch.log(sim_matrix / (denom + eps) + eps) # loss matrix
|
||||
|
||||
labels = labels.contiguous().view(-1, 1)
|
||||
Mask = torch.eq(labels, labels.t()).float().to(device)
|
||||
#Mask = eye * torch.stack([labels == labels[i] for i in range(labels.size(0))]).float().to(device)
|
||||
Mask = Mask / (Mask.sum(dim=1, keepdim=True) + eps)
|
||||
|
||||
loss = torch.sum(Mask * sim_matrix) / (2 * B)
|
||||
|
||||
return loss
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
from torch.optim.lr_scheduler import _LRScheduler
|
||||
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
||||
|
||||
|
||||
class GradualWarmupScheduler(_LRScheduler):
|
||||
""" Gradually warm-up(increasing) learning rate in optimizer.
|
||||
Proposed in 'Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour'.
|
||||
|
||||
Args:
|
||||
optimizer (Optimizer): Wrapped optimizer.
|
||||
multiplier: target learning rate = base lr * multiplier if multiplier > 1.0. if multiplier = 1.0, lr starts from 0 and ends up with the base_lr.
|
||||
total_epoch: target learning rate is reached at total_epoch, gradually
|
||||
after_scheduler: after target_epoch, use this scheduler(eg. ReduceLROnPlateau)
|
||||
"""
|
||||
|
||||
def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):
|
||||
self.multiplier = multiplier
|
||||
if self.multiplier < 1.:
|
||||
raise ValueError('multiplier should be greater thant or equal to 1.')
|
||||
self.total_epoch = total_epoch
|
||||
self.after_scheduler = after_scheduler
|
||||
self.finished = False
|
||||
super(GradualWarmupScheduler, self).__init__(optimizer)
|
||||
|
||||
def get_lr(self):
|
||||
if self.last_epoch > self.total_epoch:
|
||||
if self.after_scheduler:
|
||||
if not self.finished:
|
||||
self.after_scheduler.base_lrs = [base_lr * self.multiplier for base_lr in self.base_lrs]
|
||||
self.finished = True
|
||||
return self.after_scheduler.get_lr()
|
||||
return [base_lr * self.multiplier for base_lr in self.base_lrs]
|
||||
|
||||
if self.multiplier == 1.0:
|
||||
return [base_lr * (float(self.last_epoch) / self.total_epoch) for base_lr in self.base_lrs]
|
||||
else:
|
||||
return [base_lr * ((self.multiplier - 1.) * self.last_epoch / self.total_epoch + 1.) for base_lr in self.base_lrs]
|
||||
|
||||
def step_ReduceLROnPlateau(self, metrics, epoch=None):
|
||||
if epoch is None:
|
||||
epoch = self.last_epoch + 1
|
||||
self.last_epoch = epoch if epoch != 0 else 1 # ReduceLROnPlateau is called at the end of epoch, whereas others are called at beginning
|
||||
if self.last_epoch <= self.total_epoch:
|
||||
warmup_lr = [base_lr * ((self.multiplier - 1.) * self.last_epoch / self.total_epoch + 1.) for base_lr in self.base_lrs]
|
||||
for param_group, lr in zip(self.optimizer.param_groups, warmup_lr):
|
||||
param_group['lr'] = lr
|
||||
else:
|
||||
if epoch is None:
|
||||
self.after_scheduler.step(metrics, None)
|
||||
else:
|
||||
self.after_scheduler.step(metrics, epoch - self.total_epoch)
|
||||
|
||||
def step(self, epoch=None, metrics=None):
|
||||
if type(self.after_scheduler) != ReduceLROnPlateau:
|
||||
if self.finished and self.after_scheduler:
|
||||
if epoch is None:
|
||||
self.after_scheduler.step(None)
|
||||
else:
|
||||
self.after_scheduler.step(epoch - self.total_epoch)
|
||||
else:
|
||||
return super(GradualWarmupScheduler, self).step(epoch)
|
||||
else:
|
||||
self.step_ReduceLROnPlateau(metrics, epoch)
|
||||
@@ -0,0 +1,33 @@
|
||||
def setup(mode, P):
|
||||
fname = f'{P.dataset}_{P.model}_{mode}_{P.res}'
|
||||
|
||||
if mode == 'sup_linear':
|
||||
from .sup_linear import train
|
||||
elif mode == 'sup_CSI_linear':
|
||||
from .sup_CSI_linear import train
|
||||
elif mode == 'sup_simclr':
|
||||
from .sup_simclr import train
|
||||
elif mode == 'sup_simclr_CSI':
|
||||
assert P.batch_size == 32
|
||||
# currently only support rotation
|
||||
from .sup_simclr_CSI import train
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
if P.suffix is not None:
|
||||
fname += f'_{P.suffix}'
|
||||
|
||||
return train, fname
|
||||
|
||||
|
||||
def update_comp_loss(loss_dict, loss_in, loss_out, loss_diff, batch_size):
|
||||
loss_dict['pos'].update(loss_in, batch_size)
|
||||
loss_dict['neg'].update(loss_out, batch_size)
|
||||
loss_dict['diff'].update(loss_diff, batch_size)
|
||||
|
||||
|
||||
def summary_comp_loss(logger, tag, loss_dict, epoch):
|
||||
logger.scalar_summary(f'{tag}/pos', loss_dict['pos'].average, epoch)
|
||||
logger.scalar_summary(f'{tag}/neg', loss_dict['neg'].average, epoch)
|
||||
logger.scalar_summary(f'{tag}', loss_dict['diff'].average, epoch)
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,130 @@
|
||||
import time
|
||||
|
||||
import torch.optim
|
||||
import torch.optim.lr_scheduler as lr_scheduler
|
||||
|
||||
import models.transform_layers as TL
|
||||
from utils.utils import AverageMeter, normalize
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
hflip = TL.HorizontalFlipLayer().to(device)
|
||||
|
||||
|
||||
def train(P, epoch, model, criterion, optimizer, scheduler, loader, logger=None,
|
||||
simclr_aug=None, linear=None, linear_optim=None):
|
||||
|
||||
if P.multi_gpu:
|
||||
rotation_linear = model.module.shift_cls_layer
|
||||
joint_linear = model.module.joint_distribution_layer
|
||||
else:
|
||||
rotation_linear = model.shift_cls_layer
|
||||
joint_linear = model.joint_distribution_layer
|
||||
|
||||
if epoch == 1:
|
||||
# define optimizer and save in P (argument)
|
||||
milestones = [int(0.6 * P.epochs), int(0.75 * P.epochs), int(0.9 * P.epochs)]
|
||||
|
||||
linear_optim = torch.optim.SGD(linear.parameters(),
|
||||
lr=1e-1, weight_decay=P.weight_decay)
|
||||
P.linear_optim = linear_optim
|
||||
P.linear_scheduler = lr_scheduler.MultiStepLR(P.linear_optim, gamma=0.1, milestones=milestones)
|
||||
|
||||
rotation_linear_optim = torch.optim.SGD(rotation_linear.parameters(),
|
||||
lr=1e-1, weight_decay=P.weight_decay)
|
||||
P.rotation_linear_optim = rotation_linear_optim
|
||||
P.rot_scheduler = lr_scheduler.MultiStepLR(P.rotation_linear_optim, gamma=0.1, milestones=milestones)
|
||||
|
||||
joint_linear_optim = torch.optim.SGD(joint_linear.parameters(),
|
||||
lr=1e-1, weight_decay=P.weight_decay)
|
||||
P.joint_linear_optim = joint_linear_optim
|
||||
P.joint_scheduler = lr_scheduler.MultiStepLR(P.joint_linear_optim, gamma=0.1, milestones=milestones)
|
||||
|
||||
if logger is None:
|
||||
log_ = print
|
||||
else:
|
||||
log_ = logger.log
|
||||
|
||||
batch_time = AverageMeter()
|
||||
data_time = AverageMeter()
|
||||
|
||||
losses = dict()
|
||||
losses['cls'] = AverageMeter()
|
||||
losses['rot'] = AverageMeter()
|
||||
|
||||
check = time.time()
|
||||
for n, (images, labels) in enumerate(loader):
|
||||
model.eval()
|
||||
count = n * P.n_gpus # number of trained samples
|
||||
|
||||
data_time.update(time.time() - check)
|
||||
check = time.time()
|
||||
|
||||
### SimCLR loss ###
|
||||
if P.dataset != 'imagenet':
|
||||
batch_size = images.size(0)
|
||||
images = images.to(device)
|
||||
images = hflip(images) # 2B with hflip
|
||||
else:
|
||||
batch_size = images[0].size(0)
|
||||
images = images[0].to(device)
|
||||
|
||||
labels = labels.to(device)
|
||||
images = torch.cat([torch.rot90(images, rot, (2, 3)) for rot in range(4)]) # 4B
|
||||
rot_labels = torch.cat([torch.ones_like(labels) * k for k in range(4)], 0) # B -> 4B
|
||||
joint_labels = torch.cat([labels + P.n_classes * i for i in range(4)], dim=0)
|
||||
|
||||
images = simclr_aug(images) # simclr augmentation
|
||||
_, outputs_aux = model(images, penultimate=True)
|
||||
penultimate = outputs_aux['penultimate'].detach()
|
||||
|
||||
outputs = linear(penultimate[0:batch_size]) # only use 0 degree samples for linear eval
|
||||
outputs_rot = rotation_linear(penultimate)
|
||||
outputs_joint = joint_linear(penultimate)
|
||||
|
||||
loss_ce = criterion(outputs, labels)
|
||||
loss_rot = criterion(outputs_rot, rot_labels)
|
||||
loss_joint = criterion(outputs_joint, joint_labels)
|
||||
|
||||
### CE loss ###
|
||||
P.linear_optim.zero_grad()
|
||||
loss_ce.backward()
|
||||
P.linear_optim.step()
|
||||
|
||||
### Rot loss ###
|
||||
P.rotation_linear_optim.zero_grad()
|
||||
loss_rot.backward()
|
||||
P.rotation_linear_optim.step()
|
||||
|
||||
### Joint loss ###
|
||||
P.joint_linear_optim.zero_grad()
|
||||
loss_joint.backward()
|
||||
P.joint_linear_optim.step()
|
||||
|
||||
### optimizer learning rate ###
|
||||
lr = P.linear_optim.param_groups[0]['lr']
|
||||
|
||||
batch_time.update(time.time() - check)
|
||||
|
||||
### Log losses ###
|
||||
losses['cls'].update(loss_ce.item(), batch_size)
|
||||
losses['rot'].update(loss_rot.item(), batch_size)
|
||||
|
||||
if count % 50 == 0:
|
||||
log_('[Epoch %3d; %3d] [Time %.3f] [Data %.3f] [LR %.5f]\n'
|
||||
'[LossC %f] [LossR %f]' %
|
||||
(epoch, count, batch_time.value, data_time.value, lr,
|
||||
losses['cls'].value, losses['rot'].value))
|
||||
check = time.time()
|
||||
|
||||
P.linear_scheduler.step()
|
||||
P.rot_scheduler.step()
|
||||
P.joint_scheduler.step()
|
||||
|
||||
log_('[DONE] [Time %.3f] [Data %.3f] [LossC %f] [LossR %f]' %
|
||||
(batch_time.average, data_time.average,
|
||||
losses['cls'].average, losses['rot'].average))
|
||||
|
||||
if logger is not None:
|
||||
logger.scalar_summary('train/loss_cls', losses['cls'].average, epoch)
|
||||
logger.scalar_summary('train/loss_rot', losses['rot'].average, epoch)
|
||||
logger.scalar_summary('train/batch_time', batch_time.average, epoch)
|
||||
@@ -0,0 +1,91 @@
|
||||
import time
|
||||
|
||||
import torch.optim
|
||||
import torch.optim.lr_scheduler as lr_scheduler
|
||||
|
||||
import models.transform_layers as TL
|
||||
from utils.utils import AverageMeter, normalize
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
hflip = TL.HorizontalFlipLayer().to(device)
|
||||
|
||||
|
||||
def train(P, epoch, model, criterion, optimizer, scheduler, loader, logger=None,
|
||||
simclr_aug=None, linear=None, linear_optim=None):
|
||||
|
||||
if epoch == 1:
|
||||
# define optimizer and save in P (argument)
|
||||
milestones = [int(0.6 * P.epochs), int(0.75 * P.epochs), int(0.9 * P.epochs)]
|
||||
|
||||
linear_optim = torch.optim.SGD(linear.parameters(),
|
||||
lr=1e-1, weight_decay=P.weight_decay)
|
||||
P.linear_optim = linear_optim
|
||||
P.linear_scheduler = lr_scheduler.MultiStepLR(P.linear_optim, gamma=0.1, milestones=milestones)
|
||||
|
||||
if logger is None:
|
||||
log_ = print
|
||||
else:
|
||||
log_ = logger.log
|
||||
|
||||
batch_time = AverageMeter()
|
||||
data_time = AverageMeter()
|
||||
|
||||
losses = dict()
|
||||
losses['cls'] = AverageMeter()
|
||||
|
||||
check = time.time()
|
||||
for n, (images, labels) in enumerate(loader):
|
||||
model.eval()
|
||||
count = n * P.n_gpus # number of trained samples
|
||||
|
||||
data_time.update(time.time() - check)
|
||||
check = time.time()
|
||||
|
||||
### SimCLR loss ###
|
||||
if P.dataset != 'imagenet':
|
||||
batch_size = images.size(0)
|
||||
images = images.to(device)
|
||||
images = hflip(images) # 2B with hflip
|
||||
else:
|
||||
batch_size = images[0].size(0)
|
||||
images = images[0].to(device)
|
||||
|
||||
labels = labels.to(device)
|
||||
|
||||
images = simclr_aug(images) # simclr augmentation
|
||||
_, outputs_aux = model(images, penultimate=True)
|
||||
penultimate = outputs_aux['penultimate'].detach()
|
||||
|
||||
outputs = linear(penultimate[0:batch_size]) # only use 0 degree samples for linear eval
|
||||
|
||||
loss_ce = criterion(outputs, labels)
|
||||
|
||||
### CE loss ###
|
||||
P.linear_optim.zero_grad()
|
||||
loss_ce.backward()
|
||||
P.linear_optim.step()
|
||||
|
||||
### optimizer learning rate ###
|
||||
lr = P.linear_optim.param_groups[0]['lr']
|
||||
|
||||
batch_time.update(time.time() - check)
|
||||
|
||||
### Log losses ###
|
||||
losses['cls'].update(loss_ce.item(), batch_size)
|
||||
|
||||
if count % 50 == 0:
|
||||
log_('[Epoch %3d; %3d] [Time %.3f] [Data %.3f] [LR %.5f]\n'
|
||||
'[LossC %f]' %
|
||||
(epoch, count, batch_time.value, data_time.value, lr,
|
||||
losses['cls'].value, ))
|
||||
check = time.time()
|
||||
|
||||
P.linear_scheduler.step()
|
||||
|
||||
log_('[DONE] [Time %.3f] [Data %.3f] [LossC %f]' %
|
||||
(batch_time.average, data_time.average,
|
||||
losses['cls'].average))
|
||||
|
||||
if logger is not None:
|
||||
logger.scalar_summary('train/loss_cls', losses['cls'].average, epoch)
|
||||
logger.scalar_summary('train/batch_time', batch_time.average, epoch)
|
||||
@@ -0,0 +1,104 @@
|
||||
import time
|
||||
|
||||
import torch.optim
|
||||
|
||||
import models.transform_layers as TL
|
||||
from training.contrastive_loss import get_similarity_matrix, Supervised_NT_xent
|
||||
from utils.utils import AverageMeter, normalize
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
hflip = TL.HorizontalFlipLayer().to(device)
|
||||
|
||||
|
||||
def train(P, epoch, model, criterion, optimizer, scheduler, loader, logger=None,
|
||||
simclr_aug=None, linear=None, linear_optim=None):
|
||||
|
||||
assert simclr_aug is not None
|
||||
assert P.sim_lambda == 1.0
|
||||
|
||||
if logger is None:
|
||||
log_ = print
|
||||
else:
|
||||
log_ = logger.log
|
||||
|
||||
batch_time = AverageMeter()
|
||||
data_time = AverageMeter()
|
||||
|
||||
losses = dict()
|
||||
losses['cls'] = AverageMeter()
|
||||
losses['sim'] = AverageMeter()
|
||||
losses['simnorm'] = AverageMeter()
|
||||
|
||||
check = time.time()
|
||||
for n, (images, labels) in enumerate(loader):
|
||||
model.train()
|
||||
count = n * P.n_gpus # number of trained samples
|
||||
|
||||
data_time.update(time.time() - check)
|
||||
check = time.time()
|
||||
|
||||
### SimCLR loss ###
|
||||
if P.dataset != 'imagenet' and P.dataset != 'CNMC' and P.dataset != 'CNMC_grayscale':
|
||||
batch_size = images.size(0)
|
||||
images = images.to(device)
|
||||
images_pair = hflip(images.repeat(2, 1, 1, 1)) # 2B with hflip
|
||||
else:
|
||||
batch_size = images[0].size(0)
|
||||
images1, images2 = images[0].to(device), images[1].to(device)
|
||||
images_pair = torch.cat([images1, images2], dim=0) # 2B
|
||||
|
||||
labels = labels.to(device)
|
||||
|
||||
images_pair = simclr_aug(images_pair) # simclr augmentation
|
||||
|
||||
_, outputs_aux = model(images_pair, simclr=True, penultimate=True)
|
||||
|
||||
simclr = normalize(outputs_aux['simclr']) # normalize
|
||||
sim_matrix = get_similarity_matrix(simclr, multi_gpu=P.multi_gpu)
|
||||
loss_sim = Supervised_NT_xent(sim_matrix, labels=labels, temperature=0.07, multi_gpu=P.multi_gpu) * P.sim_lambda
|
||||
|
||||
### total loss ###
|
||||
loss = loss_sim
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
scheduler.step(epoch - 1 + n / len(loader))
|
||||
lr = optimizer.param_groups[0]['lr']
|
||||
|
||||
batch_time.update(time.time() - check)
|
||||
|
||||
### Post-processing stuffs ###
|
||||
simclr_norm = outputs_aux['simclr'].norm(dim=1).mean()
|
||||
|
||||
### Linear evaluation ###
|
||||
outputs_linear_eval = linear(outputs_aux['penultimate'].detach())
|
||||
loss_linear = criterion(outputs_linear_eval, labels.repeat(2))
|
||||
|
||||
linear_optim.zero_grad()
|
||||
loss_linear.backward()
|
||||
linear_optim.step()
|
||||
|
||||
### Log losses ###
|
||||
losses['cls'].update(0, batch_size)
|
||||
losses['sim'].update(loss_sim.item(), batch_size)
|
||||
losses['simnorm'].update(simclr_norm.item(), batch_size)
|
||||
|
||||
if count % 50 == 0:
|
||||
log_('[Epoch %3d; %3d] [Time %.3f] [Data %.3f] [LR %.5f]\n'
|
||||
'[LossC %f] [LossSim %f] [SimNorm %f]' %
|
||||
(epoch, count, batch_time.value, data_time.value, lr,
|
||||
losses['cls'].value, losses['sim'].value, losses['simnorm'].value))
|
||||
|
||||
check = time.time()
|
||||
|
||||
log_('[DONE] [Time %.3f] [Data %.3f] [LossC %f] [LossSim %f] [SimNorm %f]' %
|
||||
(batch_time.average, data_time.average,
|
||||
losses['cls'].average, losses['sim'].average, losses['simnorm'].average))
|
||||
|
||||
if logger is not None:
|
||||
logger.scalar_summary('train/loss_cls', losses['cls'].average, epoch)
|
||||
logger.scalar_summary('train/loss_sim', losses['sim'].average, epoch)
|
||||
logger.scalar_summary('train/batch_time', batch_time.average, epoch)
|
||||
logger.scalar_summary('train/simclr_norm', losses['simnorm'].average, epoch)
|
||||
@@ -0,0 +1,111 @@
|
||||
import time
|
||||
|
||||
import torch.optim
|
||||
|
||||
import models.transform_layers as TL
|
||||
from training.contrastive_loss import get_similarity_matrix, Supervised_NT_xent
|
||||
from utils.utils import AverageMeter, normalize
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
hflip = TL.HorizontalFlipLayer().to(device)
|
||||
|
||||
|
||||
def train(P, epoch, model, criterion, optimizer, scheduler, loader, logger=None,
|
||||
simclr_aug=None, linear=None, linear_optim=None):
|
||||
|
||||
# currently only support rotation shifting augmentation
|
||||
assert simclr_aug is not None
|
||||
assert P.sim_lambda == 1.0
|
||||
|
||||
if logger is None:
|
||||
log_ = print
|
||||
else:
|
||||
log_ = logger.log
|
||||
|
||||
batch_time = AverageMeter()
|
||||
data_time = AverageMeter()
|
||||
|
||||
losses = dict()
|
||||
losses['cls'] = AverageMeter()
|
||||
losses['sim'] = AverageMeter()
|
||||
|
||||
check = time.time()
|
||||
for n, (images, labels) in enumerate(loader):
|
||||
model.train()
|
||||
count = n * P.n_gpus # number of trained samples
|
||||
|
||||
data_time.update(time.time() - check)
|
||||
check = time.time()
|
||||
|
||||
### SimCLR loss ###
|
||||
if P.dataset != 'imagenet' and P.dataset != 'CNMC' and P.dataset != 'CNMC_grayscale':
|
||||
batch_size = images.size(0)
|
||||
images = images.to(device)
|
||||
images1, images2 = hflip(images.repeat(2, 1, 1, 1)).chunk(2) # hflip
|
||||
else:
|
||||
batch_size = images[0].size(0)
|
||||
images1, images2 = images[0].to(device), images[1].to(device)
|
||||
#print("\nImages" + str(images.shape) + "\n")
|
||||
|
||||
images1 = torch.cat([torch.rot90(images1, rot, (2, 3)) for rot in range(4)]) # 4B
|
||||
images2 = torch.cat([torch.rot90(images2, rot, (2, 3)) for rot in range(4)]) # 4B
|
||||
images_pair = torch.cat([images1, images2], dim=0) # 8B
|
||||
|
||||
labels = labels.to(device)
|
||||
rot_sim_labels = torch.cat([labels + P.n_classes * i for i in range(4)], dim=0)
|
||||
rot_sim_labels = rot_sim_labels.to(device)
|
||||
|
||||
images_pair = simclr_aug(images_pair) # simclr augment
|
||||
_, outputs_aux = model(images_pair, simclr=True, penultimate=True)
|
||||
|
||||
simclr = normalize(outputs_aux['simclr']) # normalize
|
||||
sim_matrix = get_similarity_matrix(simclr, multi_gpu=P.multi_gpu)
|
||||
loss_sim = Supervised_NT_xent(sim_matrix, labels=rot_sim_labels,
|
||||
temperature=0.07, multi_gpu=P.multi_gpu) * P.sim_lambda
|
||||
|
||||
### total loss ###
|
||||
loss = loss_sim
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
scheduler.step(epoch - 1 + n / len(loader))
|
||||
lr = optimizer.param_groups[0]['lr']
|
||||
|
||||
batch_time.update(time.time() - check)
|
||||
|
||||
### Post-processing stuffs ###
|
||||
penul_1 = outputs_aux['penultimate'][:batch_size]
|
||||
penul_2 = outputs_aux['penultimate'][4 * batch_size: 5 * batch_size]
|
||||
outputs_aux['penultimate'] = torch.cat([penul_1, penul_2]) # only use original rotation
|
||||
|
||||
### Linear evaluation ###
|
||||
outputs_linear_eval = linear(outputs_aux['penultimate'].detach())
|
||||
loss_linear = criterion(outputs_linear_eval, labels.repeat(2))
|
||||
|
||||
linear_optim.zero_grad()
|
||||
loss_linear.backward()
|
||||
linear_optim.step()
|
||||
|
||||
### Log losses ###
|
||||
losses['cls'].update(0, batch_size)
|
||||
losses['sim'].update(loss_sim.item(), batch_size)
|
||||
|
||||
if count % 50 == 0:
|
||||
log_('[Epoch %3d; %3d] [Time %.3f] [Data %.3f] [LR %.5f]\n'
|
||||
'[LossC %f] [LossSim %f]' %
|
||||
(epoch, count, batch_time.value, data_time.value, lr,
|
||||
losses['cls'].value, losses['sim'].value))
|
||||
|
||||
check = time.time()
|
||||
|
||||
log_('[DONE] [Time %.3f] [Data %.3f] [LossC %f] [LossSim %f]' %
|
||||
(batch_time.average, data_time.average,
|
||||
losses['cls'].average, losses['sim'].average))
|
||||
|
||||
if logger is not None:
|
||||
logger.scalar_summary('train/loss_cls', losses['cls'].average, epoch)
|
||||
logger.scalar_summary('train/loss_sim', losses['sim'].average, epoch)
|
||||
logger.scalar_summary('train/batch_time', batch_time.average, epoch)
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
def setup(mode, P):
|
||||
fname = f'{P.dataset}_{P.model}_unsup_{mode}_{P.res}'
|
||||
|
||||
if mode == 'simclr':
|
||||
from .simclr import train
|
||||
elif mode == 'simclr_CSI':
|
||||
from .simclr_CSI import train
|
||||
fname += f'_shift_{P.shift_trans_type}_resize_factor{P.resize_factor}_color_dist{P.color_distort}'
|
||||
if P.shift_trans_type == 'gauss':
|
||||
fname += f'_gauss_sigma{P.gauss_sigma}'
|
||||
elif P.shift_trans_type == 'randpers':
|
||||
fname += f'_distortion_scale{P.distortion_scale}'
|
||||
elif P.shift_trans_type == 'sharp':
|
||||
fname += f'_sharpness_factor{P.sharpness_factor}'
|
||||
elif P.shift_trans_type == 'sharp':
|
||||
fname += f'_nmean_{P.noise_mean}_nstd_{P.noise_std}'
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
if P.one_class_idx is not None:
|
||||
fname += f'_one_class_{P.one_class_idx}'
|
||||
|
||||
if P.suffix is not None:
|
||||
fname += f'_{P.suffix}'
|
||||
|
||||
return train, fname
|
||||
|
||||
|
||||
def update_comp_loss(loss_dict, loss_in, loss_out, loss_diff, batch_size):
|
||||
loss_dict['pos'].update(loss_in, batch_size)
|
||||
loss_dict['neg'].update(loss_out, batch_size)
|
||||
loss_dict['diff'].update(loss_diff, batch_size)
|
||||
|
||||
|
||||
def summary_comp_loss(logger, tag, loss_dict, epoch):
|
||||
logger.scalar_summary(f'{tag}/pos', loss_dict['pos'].average, epoch)
|
||||
logger.scalar_summary(f'{tag}/neg', loss_dict['neg'].average, epoch)
|
||||
logger.scalar_summary(f'{tag}', loss_dict['diff'].average, epoch)
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,101 @@
|
||||
import time
|
||||
|
||||
import torch.optim
|
||||
|
||||
import models.transform_layers as TL
|
||||
from training.contrastive_loss import get_similarity_matrix, NT_xent
|
||||
from utils.utils import AverageMeter, normalize
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
hflip = TL.HorizontalFlipLayer().to(device)
|
||||
|
||||
|
||||
def train(P, epoch, model, criterion, optimizer, scheduler, loader, logger=None,
|
||||
simclr_aug=None, linear=None, linear_optim=None):
|
||||
|
||||
assert simclr_aug is not None
|
||||
assert P.sim_lambda == 1.0
|
||||
|
||||
if logger is None:
|
||||
log_ = print
|
||||
else:
|
||||
log_ = logger.log
|
||||
|
||||
batch_time = AverageMeter()
|
||||
data_time = AverageMeter()
|
||||
|
||||
losses = dict()
|
||||
losses['cls'] = AverageMeter()
|
||||
losses['sim'] = AverageMeter()
|
||||
|
||||
check = time.time()
|
||||
for n, (images, labels) in enumerate(loader):
|
||||
model.train()
|
||||
count = n * P.n_gpus # number of trained samples
|
||||
|
||||
data_time.update(time.time() - check)
|
||||
check = time.time()
|
||||
|
||||
### SimCLR loss ###
|
||||
if P.dataset != 'imagenet':
|
||||
batch_size = images.size(0)
|
||||
images = images.to(device)
|
||||
images_pair = hflip(images.repeat(2, 1, 1, 1)) # 2B with hflip
|
||||
else:
|
||||
batch_size = images[0].size(0)
|
||||
images1, images2 = images[0].to(device), images[1].to(device)
|
||||
images_pair = torch.cat([images1, images2], dim=0) # 2B
|
||||
|
||||
labels = labels.to(device)
|
||||
|
||||
images_pair = simclr_aug(images_pair) # transform
|
||||
|
||||
_, outputs_aux = model(images_pair, simclr=True, penultimate=True)
|
||||
|
||||
simclr = normalize(outputs_aux['simclr']) # normalize
|
||||
sim_matrix = get_similarity_matrix(simclr, multi_gpu=P.multi_gpu)
|
||||
loss_sim = NT_xent(sim_matrix, temperature=0.5) * P.sim_lambda
|
||||
|
||||
### total loss ###
|
||||
loss = loss_sim
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
scheduler.step(epoch - 1 + n / len(loader))
|
||||
lr = optimizer.param_groups[0]['lr']
|
||||
|
||||
batch_time.update(time.time() - check)
|
||||
|
||||
### Post-processing stuffs ###
|
||||
simclr_norm = outputs_aux['simclr'].norm(dim=1).mean()
|
||||
|
||||
### Linear evaluation ###
|
||||
outputs_linear_eval = linear(outputs_aux['penultimate'].detach())
|
||||
loss_linear = criterion(outputs_linear_eval, labels.repeat(2))
|
||||
|
||||
linear_optim.zero_grad()
|
||||
loss_linear.backward()
|
||||
linear_optim.step()
|
||||
|
||||
### Log losses ###
|
||||
losses['cls'].update(0, batch_size)
|
||||
losses['sim'].update(loss_sim.item(), batch_size)
|
||||
|
||||
if count % 50 == 0:
|
||||
log_('[Epoch %3d; %3d] [Time %.3f] [Data %.3f] [LR %.5f]\n'
|
||||
'[LossC %f] [LossSim %f]' %
|
||||
(epoch, count, batch_time.value, data_time.value, lr,
|
||||
losses['cls'].value, losses['sim'].value))
|
||||
|
||||
check = time.time()
|
||||
|
||||
log_('[DONE] [Time %.3f] [Data %.3f] [LossC %f] [LossSim %f]' %
|
||||
(batch_time.average, data_time.average,
|
||||
losses['cls'].average, losses['sim'].average))
|
||||
|
||||
if logger is not None:
|
||||
logger.scalar_summary('train/loss_cls', losses['cls'].average, epoch)
|
||||
logger.scalar_summary('train/loss_sim', losses['sim'].average, epoch)
|
||||
logger.scalar_summary('train/batch_time', batch_time.average, epoch)
|
||||
@@ -0,0 +1,114 @@
|
||||
import time
|
||||
|
||||
import torch.optim
|
||||
|
||||
import models.transform_layers as TL
|
||||
from training.contrastive_loss import get_similarity_matrix, NT_xent
|
||||
from utils.utils import AverageMeter, normalize
|
||||
|
||||
device = torch.device(f"cuda" if torch.cuda.is_available() else "cpu")
|
||||
hflip = TL.HorizontalFlipLayer().to(device)
|
||||
|
||||
|
||||
def train(P, epoch, model, criterion, optimizer, scheduler, loader, logger=None,
|
||||
simclr_aug=None, linear=None, linear_optim=None):
|
||||
|
||||
assert simclr_aug is not None
|
||||
assert P.sim_lambda == 1.0 # to avoid mistake
|
||||
assert P.K_shift > 1
|
||||
|
||||
if logger is None:
|
||||
log_ = print
|
||||
else:
|
||||
log_ = logger.log
|
||||
|
||||
batch_time = AverageMeter()
|
||||
data_time = AverageMeter()
|
||||
|
||||
losses = dict()
|
||||
losses['cls'] = AverageMeter()
|
||||
losses['sim'] = AverageMeter()
|
||||
losses['shift'] = AverageMeter()
|
||||
|
||||
check = time.time()
|
||||
for n, (images, labels) in enumerate(loader):
|
||||
model.train()
|
||||
count = n * P.n_gpus # number of trained samples
|
||||
|
||||
data_time.update(time.time() - check)
|
||||
check = time.time()
|
||||
|
||||
### SimCLR loss ###
|
||||
if P.dataset != 'imagenet' and P.dataset != 'CNMC' and P.dataset != 'CNMC_grayscale':
|
||||
batch_size = images.size(0)
|
||||
images = images.to(device)
|
||||
images1, images2 = hflip(images.repeat(2, 1, 1, 1)).chunk(2) # hflip
|
||||
else:
|
||||
batch_size = images[0].size(0)
|
||||
images1, images2 = images[0].to(device), images[1].to(device)
|
||||
labels = labels.to(device)
|
||||
|
||||
images1 = torch.cat([P.shift_trans(images1, k) for k in range(P.K_shift)])
|
||||
images2 = torch.cat([P.shift_trans(images2, k) for k in range(P.K_shift)])
|
||||
|
||||
shift_labels = torch.cat([torch.ones_like(labels) * k for k in range(P.K_shift)], 0) # B -> 4B
|
||||
shift_labels = shift_labels.repeat(2)
|
||||
|
||||
images_pair = torch.cat([images1, images2], dim=0) # 8B
|
||||
images_pair = simclr_aug(images_pair) # transform
|
||||
|
||||
_, outputs_aux = model(images_pair, simclr=True, penultimate=True, shift=True)
|
||||
|
||||
simclr = normalize(outputs_aux['simclr']) # normalize
|
||||
sim_matrix = get_similarity_matrix(simclr, multi_gpu=P.multi_gpu)
|
||||
loss_sim = NT_xent(sim_matrix, temperature=0.5) * P.sim_lambda
|
||||
|
||||
loss_shift = criterion(outputs_aux['shift'], shift_labels)
|
||||
|
||||
### total loss ###
|
||||
loss = loss_sim + loss_shift
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
scheduler.step(epoch - 1 + n / len(loader))
|
||||
lr = optimizer.param_groups[0]['lr']
|
||||
|
||||
batch_time.update(time.time() - check)
|
||||
|
||||
### Post-processing stuffs ###
|
||||
simclr_norm = outputs_aux['simclr'].norm(dim=1).mean()
|
||||
|
||||
penul_1 = outputs_aux['penultimate'][:batch_size]
|
||||
penul_2 = outputs_aux['penultimate'][P.K_shift * batch_size: (P.K_shift + 1) * batch_size]
|
||||
outputs_aux['penultimate'] = torch.cat([penul_1, penul_2]) # only use original rotation
|
||||
|
||||
### Linear evaluation ###
|
||||
outputs_linear_eval = linear(outputs_aux['penultimate'].detach())
|
||||
loss_linear = criterion(outputs_linear_eval, labels.repeat(2))
|
||||
|
||||
linear_optim.zero_grad()
|
||||
loss_linear.backward()
|
||||
linear_optim.step()
|
||||
|
||||
losses['cls'].update(0, batch_size)
|
||||
losses['sim'].update(loss_sim.item(), batch_size)
|
||||
losses['shift'].update(loss_shift.item(), batch_size)
|
||||
|
||||
if count % 50 == 0:
|
||||
log_('[Epoch %3d; %3d] [Time %.3f] [Data %.3f] [LR %.5f]\n'
|
||||
'[LossC %f] [LossSim %f] [LossShift %f]' %
|
||||
(epoch, count, batch_time.value, data_time.value, lr,
|
||||
losses['cls'].value, losses['sim'].value, losses['shift'].value))
|
||||
|
||||
log_('[DONE] [Time %.3f] [Data %.3f] [LossC %f] [LossSim %f] [LossShift %f]' %
|
||||
(batch_time.average, data_time.average,
|
||||
losses['cls'].average, losses['sim'].average, losses['shift'].average))
|
||||
|
||||
if logger is not None:
|
||||
logger.scalar_summary('train/loss_cls', losses['cls'].average, epoch)
|
||||
logger.scalar_summary('train/loss_sim', losses['sim'].average, epoch)
|
||||
logger.scalar_summary('train/loss_shift', losses['shift'].average, epoch)
|
||||
logger.scalar_summary('train/batch_time', batch_time.average, epoch)
|
||||
|
||||
Reference in New Issue
Block a user