init
This commit is contained in:
+119
@@ -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)
|
||||
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.
@@ -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()
|
||||
@@ -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
@@ -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)
|
||||
Reference in New Issue
Block a user