This commit is contained in:
2022-04-29 19:26:47 +02:00
commit d1ce7b933f
110 changed files with 17469 additions and 0 deletions
+119
View File
@@ -0,0 +1,119 @@
"""
References:
- https://github.com/PyTorchLightning/PyTorch-Lightning-Bolts/blob/master/pl_bolts/optimizers/lars_scheduling.py
- https://github.com/NVIDIA/apex/blob/master/apex/parallel/LARC.py
- https://arxiv.org/pdf/1708.03888.pdf
- https://github.com/noahgolmant/pytorch-lars/blob/master/lars.py
"""
import torch
from .wrapper import OptimWrapper
# from torchlars._adaptive_lr import compute_adaptive_lr # Impossible to build extensions
__all__ = ["LARS"]
class LARS(OptimWrapper):
"""Implements 'LARS (Layer-wise Adaptive Rate Scaling)'__ as Optimizer a
:class:`~torch.optim.Optimizer` wrapper.
__ : https://arxiv.org/abs/1708.03888
Wraps an arbitrary optimizer like :class:`torch.optim.SGD` to use LARS. If
you want to the same performance obtained with small-batch training when
you use large-batch training, LARS will be helpful::
Args:
optimizer (Optimizer):
optimizer to wrap
eps (float, optional):
epsilon to help with numerical stability while calculating the
adaptive learning rate
trust_coef (float, optional):
trust coefficient for calculating the adaptive learning rate
Example::
base_optimizer = optim.SGD(model.parameters(), lr=0.1)
optimizer = LARS(optimizer=base_optimizer)
output = model(input)
loss = loss_fn(output, target)
loss.backward()
optimizer.step()
"""
def __init__(self, optimizer, trust_coef=0.02, clip=True, eps=1e-8):
if eps < 0.0:
raise ValueError("invalid epsilon value: , %f" % eps)
if trust_coef < 0.0:
raise ValueError("invalid trust coefficient: %f" % trust_coef)
self.optim = optimizer
self.eps = eps
self.trust_coef = trust_coef
self.clip = clip
def __getstate__(self):
self.optim.__get
lars_dict = {}
lars_dict["trust_coef"] = self.trust_coef
lars_dict["clip"] = self.clip
lars_dict["eps"] = self.eps
return (self.optim, lars_dict)
def __setstate__(self, state):
self.optim, lars_dict = state
self.trust_coef = lars_dict["trust_coef"]
self.clip = lars_dict["clip"]
self.eps = lars_dict["eps"]
@torch.no_grad()
def step(self, closure=None):
weight_decays = []
for group in self.optim.param_groups:
weight_decay = group.get("weight_decay", 0)
weight_decays.append(weight_decay)
# reset weight decay
group["weight_decay"] = 0
# update the parameters
for p in group["params"]:
if p.grad is not None:
self.update_p(p, group, weight_decay)
# update the optimizer
self.optim.step(closure=closure)
# return weight decay control to optimizer
for group_idx, group in enumerate(self.optim.param_groups):
group["weight_decay"] = weight_decays[group_idx]
def update_p(self, p, group, weight_decay):
# calculate new norms
p_norm = torch.norm(p.data)
g_norm = torch.norm(p.grad.data)
if p_norm != 0 and g_norm != 0:
# calculate new lr
divisor = g_norm + p_norm * weight_decay + self.eps
adaptive_lr = (self.trust_coef * p_norm) / divisor
# clip lr
if self.clip:
adaptive_lr = min(adaptive_lr / group["lr"], 1)
# update params with clipped lr
p.grad.data += weight_decay * p.data
p.grad.data *= adaptive_lr
from torch.optim import SGD
from pylot.util import delegates, separate_kwargs
class SGDLARS(LARS):
@delegates(to=LARS.__init__)
@delegates(to=SGD.__init__, keep=True, but=["eps", "trust_coef"])
def __init__(self, params, **kwargs):
sgd_kwargs, lars_kwargs = separate_kwargs(kwargs, SGD.__init__)
optim = SGD(params, **sgd_kwargs)
super().__init__(optim, **lars_kwargs)
View File
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+114
View File
@@ -0,0 +1,114 @@
from argparse import ArgumentParser
def parse_args(default=False):
"""Command-line argument parser for training."""
parser = ArgumentParser(description='Pytorch implementation of CSI')
parser.add_argument('--dataset', help='Dataset',
choices=['cifar10', 'cifar100', 'imagenet', 'CNMC', 'CNMC_grayscale'], type=str)
parser.add_argument('--one_class_idx', help='None: multi-class, Not None: one-class',
default=None, type=int)
parser.add_argument('--model', help='Model',
choices=['resnet18', 'resnet18_imagenet'], type=str)
parser.add_argument('--mode', help='Training mode',
default='simclr', type=str)
parser.add_argument('--simclr_dim', help='Dimension of simclr layer',
default=128, type=int)
parser.add_argument('--shift_trans_type', help='shifting transformation type', default='none',
choices=['rotation', 'cutperm', 'blur', 'randpers', 'sharp', 'blur_randpers',
'blur_sharp', 'randpers_sharp', 'blur_randpers_sharp', 'noise', 'none'], type=str)
parser.add_argument("--local_rank", type=int,
default=0, help='Local rank for distributed learning')
parser.add_argument('--resume_path', help='Path to the resume checkpoint',
default=None, type=str)
parser.add_argument('--load_path', help='Path to the loading checkpoint',
default=None, type=str)
parser.add_argument("--no_strict", help='Do not strictly load state_dicts',
action='store_true')
parser.add_argument('--suffix', help='Suffix for the log dir',
default=None, type=str)
parser.add_argument('--error_step', help='Epoch steps to compute errors',
default=5, type=int)
parser.add_argument('--save_step', help='Epoch steps to save models',
default=10, type=int)
##### Training Configurations #####
parser.add_argument('--epochs', help='Epochs',
default=1000, type=int)
parser.add_argument('--optimizer', help='Optimizer',
choices=['sgd', 'lars'],
default='lars', type=str)
parser.add_argument('--lr_scheduler', help='Learning rate scheduler',
choices=['step_decay', 'cosine'],
default='cosine', type=str)
parser.add_argument('--warmup', help='Warm-up epochs',
default=10, type=int)
parser.add_argument('--lr_init', help='Initial learning rate',
default=1e-1, type=float)
parser.add_argument('--weight_decay', help='Weight decay',
default=1e-6, type=float)
parser.add_argument('--batch_size', help='Batch size',
default=128, type=int)
parser.add_argument('--test_batch_size', help='Batch size for test loader',
default=100, type=int)
parser.add_argument('--blur_sigma', help='Distortion grade',
default=2.0, type=float)
parser.add_argument('--color_distort', help='Color distortion grade',
default=0.5, type=float)
parser.add_argument('--distortion_scale', help='Perspective distortion grade',
default=0.6, type=float)
parser.add_argument('--sharpness_factor', help='Sharpening or blurring factor of image. '
'Can be any non negative number. 0 gives a blurred image, '
'1 gives the original image while 2 increases the sharpness '
'by a factor of 2.',
default=2, type=float)
parser.add_argument('--noise_mean', help='mean',
default=0, type=float)
parser.add_argument('--noise_std', help='std',
default=0.3, type=float)
##### Objective Configurations #####
parser.add_argument('--sim_lambda', help='Weight for SimCLR loss',
default=1.0, type=float)
parser.add_argument('--temperature', help='Temperature for similarity',
default=0.5, type=float)
##### Evaluation Configurations #####
parser.add_argument("--ood_dataset", help='Datasets for OOD detection',
default=None, nargs="*", type=str)
parser.add_argument("--ood_score", help='score function for OOD detection',
default=['norm_mean'], nargs="+", type=str)
parser.add_argument("--ood_layer", help='layer for OOD scores',
choices=['penultimate', 'simclr', 'shift'],
default=['simclr', 'shift'], nargs="+", type=str)
parser.add_argument("--ood_samples", help='number of samples to compute OOD score',
default=1, type=int)
parser.add_argument("--ood_batch_size", help='batch size to compute OOD score',
default=100, type=int)
parser.add_argument("--resize_factor", help='resize scale is sampled from [resize_factor, 1.0]',
default=0.08, type=float)
parser.add_argument("--resize_fix", help='resize scale is fixed to resize_factor (not (resize_factor, 1.0])',
action='store_true')
parser.add_argument("--print_score", help='print quantiles of ood score',
action='store_true')
parser.add_argument("--save_score", help='save ood score for plotting histogram',
action='store_true')
##### Process configuration option #####
parser.add_argument("--proc_step", help='choose process to initiate.',
choices=['eval', 'train'],
default=None, type=str)
parser.add_argument("--res", help='resolution of dataset',
default="32px", type=str)
if default:
return parser.parse_args('') # empty string
else:
return parser.parse_args()
+81
View File
@@ -0,0 +1,81 @@
from copy import deepcopy
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from common.common import parse_args
import models.classifier as C
from datasets import get_dataset, get_superclass_list, get_subclass_dataset
P = parse_args()
### Set torch device ###
P.n_gpus = torch.cuda.device_count()
assert P.n_gpus <= 1 # no multi GPU
P.multi_gpu = False
if torch.cuda.is_available():
torch.cuda.set_device(P.local_rank)
device = torch.device(f"cuda" if torch.cuda.is_available() else "cpu")
### Initialize dataset ###
ood_eval = P.mode == 'ood_pre'
if P.dataset == 'imagenet' and ood_eval or P.dataset == 'CNMC' and ood_eval or P.dataset == 'CNMC_grayscale' and ood_eval:
P.batch_size = 1
P.test_batch_size = 1
train_set, test_set, image_size, n_classes = get_dataset(P, dataset=P.dataset, eval=ood_eval)
P.image_size = image_size
P.n_classes = n_classes
if P.one_class_idx is not None:
cls_list = get_superclass_list(P.dataset)
P.n_superclasses = len(cls_list)
full_test_set = deepcopy(test_set) # test set of full classes
train_set = get_subclass_dataset(train_set, classes=cls_list[P.one_class_idx])
test_set = get_subclass_dataset(test_set, classes=cls_list[P.one_class_idx])
kwargs = {'pin_memory': False, 'num_workers': 2}
train_loader = DataLoader(train_set, shuffle=True, batch_size=P.batch_size, **kwargs)
test_loader = DataLoader(test_set, shuffle=False, batch_size=P.test_batch_size, **kwargs)
if P.ood_dataset is None:
if P.one_class_idx is not None:
P.ood_dataset = list(range(P.n_superclasses))
P.ood_dataset.pop(P.one_class_idx)
elif P.dataset == 'cifar10':
P.ood_dataset = ['svhn', 'lsun_resize', 'imagenet_resize', 'lsun_fix', 'imagenet_fix', 'cifar100', 'interp']
elif P.dataset == 'imagenet':
P.ood_dataset = ['cub', 'stanford_dogs', 'flowers102', 'places365', 'food_101', 'caltech_256', 'dtd', 'pets']
ood_test_loader = dict()
for ood in P.ood_dataset:
if ood == 'interp':
ood_test_loader[ood] = None # dummy loader
continue
if P.one_class_idx is not None:
ood_test_set = get_subclass_dataset(full_test_set, classes=cls_list[ood])
ood = f'one_class_{ood}' # change save name
else:
ood_test_set = get_dataset(P, dataset=ood, test_only=True, image_size=P.image_size, eval=ood_eval)
ood_test_loader[ood] = DataLoader(ood_test_set, shuffle=False, batch_size=P.test_batch_size, **kwargs)
### Initialize model ###
simclr_aug = C.get_simclr_augmentation(P, image_size=P.image_size).to(device)
P.shift_trans, P.K_shift = C.get_shift_module(P, eval=True)
P.shift_trans = P.shift_trans.to(device)
model = C.get_classifier(P.model, n_classes=P.n_classes).to(device)
model = C.get_shift_classifer(model, P.K_shift).to(device)
criterion = nn.CrossEntropyLoss().to(device)
if P.load_path is not None:
checkpoint = torch.load(P.load_path)
model.load_state_dict(checkpoint, strict=not P.no_strict)
+148
View File
@@ -0,0 +1,148 @@
from copy import deepcopy
import torch
import torch.nn as nn
import torch.optim as optim
import torch.optim.lr_scheduler as lr_scheduler
from torch.utils.data import DataLoader
from common.common import parse_args
import models.classifier as C
from datasets import get_dataset, get_superclass_list, get_subclass_dataset
from utils.utils import load_checkpoint
P = parse_args()
### Set torch device ###
if torch.cuda.is_available():
torch.cuda.set_device(P.local_rank)
device = torch.device(f"cuda" if torch.cuda.is_available() else "cpu")
P.n_gpus = torch.cuda.device_count()
if P.n_gpus > 1:
import apex
import torch.distributed as dist
from torch.utils.data.distributed import DistributedSampler
P.multi_gpu = True
torch.distributed.init_process_group(
'nccl',
init_method='env://',
world_size=P.n_gpus,
rank=P.local_rank,
)
else:
P.multi_gpu = False
### only use one ood_layer while training
P.ood_layer = P.ood_layer[0]
### Initialize dataset ###
train_set, test_set, image_size, n_classes = get_dataset(P, dataset=P.dataset)
P.image_size = image_size
P.n_classes = n_classes
if P.one_class_idx is not None:
cls_list = get_superclass_list(P.dataset)
P.n_superclasses = len(cls_list)
full_test_set = deepcopy(test_set) # test set of full classes
train_set = get_subclass_dataset(train_set, classes=cls_list[P.one_class_idx])
test_set = get_subclass_dataset(test_set, classes=cls_list[P.one_class_idx])
kwargs = {'pin_memory': False, 'num_workers': 2}
if P.multi_gpu:
train_sampler = DistributedSampler(train_set, num_replicas=P.n_gpus, rank=P.local_rank)
test_sampler = DistributedSampler(test_set, num_replicas=P.n_gpus, rank=P.local_rank)
train_loader = DataLoader(train_set, sampler=train_sampler, batch_size=P.batch_size, **kwargs)
test_loader = DataLoader(test_set, sampler=test_sampler, batch_size=P.test_batch_size, **kwargs)
else:
train_loader = DataLoader(train_set, shuffle=True, batch_size=P.batch_size, **kwargs)
test_loader = DataLoader(test_set, shuffle=False, batch_size=P.test_batch_size, **kwargs)
if P.ood_dataset is None:
if P.one_class_idx is not None:
P.ood_dataset = list(range(P.n_superclasses))
P.ood_dataset.pop(P.one_class_idx)
elif P.dataset == 'cifar10':
P.ood_dataset = ['svhn', 'lsun_resize', 'imagenet_resize', 'lsun_fix', 'imagenet_fix', 'cifar100', 'interp']
elif P.dataset == 'imagenet':
P.ood_dataset = ['cub', 'stanford_dogs', 'flowers102']
ood_test_loader = dict()
for ood in P.ood_dataset:
if ood == 'interp':
ood_test_loader[ood] = None # dummy loader
continue
if P.one_class_idx is not None:
ood_test_set = get_subclass_dataset(full_test_set, classes=cls_list[ood])
ood = f'one_class_{ood}' # change save name
else:
ood_test_set = get_dataset(P, dataset=ood, test_only=True, image_size=P.image_size)
if P.multi_gpu:
ood_sampler = DistributedSampler(ood_test_set, num_replicas=P.n_gpus, rank=P.local_rank)
ood_test_loader[ood] = DataLoader(ood_test_set, sampler=ood_sampler, batch_size=P.test_batch_size, **kwargs)
else:
ood_test_loader[ood] = DataLoader(ood_test_set, shuffle=False, batch_size=P.test_batch_size, **kwargs)
### Initialize model ###
simclr_aug = C.get_simclr_augmentation(P, image_size=P.image_size).to(device)
P.shift_trans, P.K_shift = C.get_shift_module(P, eval=True)
P.shift_trans = P.shift_trans.to(device)
model = C.get_classifier(P.model, n_classes=P.n_classes).to(device)
model = C.get_shift_classifer(model, P.K_shift).to(device)
criterion = nn.CrossEntropyLoss().to(device)
if P.optimizer == 'sgd':
optimizer = optim.SGD(model.parameters(), lr=P.lr_init, momentum=0.9, weight_decay=P.weight_decay)
lr_decay_gamma = 0.1
elif P.optimizer == 'lars':
from torchlars import LARS
base_optimizer = optim.SGD(model.parameters(), lr=P.lr_init, momentum=0.9, weight_decay=P.weight_decay)
optimizer = LARS(base_optimizer, eps=1e-8, trust_coef=0.001)
lr_decay_gamma = 0.1
else:
raise NotImplementedError()
if P.lr_scheduler == 'cosine':
scheduler = lr_scheduler.CosineAnnealingLR(optimizer, P.epochs)
elif P.lr_scheduler == 'step_decay':
milestones = [int(0.5 * P.epochs), int(0.75 * P.epochs)]
scheduler = lr_scheduler.MultiStepLR(optimizer, gamma=lr_decay_gamma, milestones=milestones)
else:
raise NotImplementedError()
from training.scheduler import GradualWarmupScheduler
scheduler_warmup = GradualWarmupScheduler(optimizer, multiplier=10.0, total_epoch=P.warmup, after_scheduler=scheduler)
if P.resume_path is not None:
resume = True
model_state, optim_state, config = load_checkpoint(P.resume_path, mode='last')
model.load_state_dict(model_state, strict=not P.no_strict)
optimizer.load_state_dict(optim_state)
start_epoch = config['epoch']
best = config['best']
error = 100.0
else:
resume = False
start_epoch = 1
best = 100.0
error = 100.0
if P.mode == 'sup_linear' or P.mode == 'sup_CSI_linear':
assert P.load_path is not None
checkpoint = torch.load(P.load_path)
model.load_state_dict(checkpoint, strict=not P.no_strict)
if P.multi_gpu:
simclr_aug = apex.parallel.DistributedDataParallel(simclr_aug, delay_allreduce=True)
model = apex.parallel.convert_syncbn_model(model)
model = apex.parallel.DistributedDataParallel(model, delay_allreduce=True)