Initial commit

This commit is contained in:
Mido Assran
2023-06-13 13:03:30 +00:00
commit 52c1ae95d0
24 changed files with 3016 additions and 0 deletions
+223
View File
@@ -0,0 +1,223 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import os
import subprocess
import time
import numpy as np
from logging import getLogger
import torch
import torchvision
_GLOBAL_SEED = 0
logger = getLogger()
def make_imagenet1k(
transform,
batch_size,
collator=None,
pin_mem=True,
num_workers=8,
world_size=1,
rank=0,
root_path=None,
image_folder=None,
training=True,
copy_data=False,
drop_last=True,
subset_file=None
):
dataset = ImageNet(
root=root_path,
image_folder=image_folder,
transform=transform,
train=training,
copy_data=copy_data,
index_targets=False)
if subset_file is not None:
dataset = ImageNetSubset(dataset, subset_file)
logger.info('ImageNet dataset created')
dist_sampler = torch.utils.data.distributed.DistributedSampler(
dataset=dataset,
num_replicas=world_size,
rank=rank)
data_loader = torch.utils.data.DataLoader(
dataset,
collate_fn=collator,
sampler=dist_sampler,
batch_size=batch_size,
drop_last=drop_last,
pin_memory=pin_mem,
num_workers=num_workers,
persistent_workers=False)
logger.info('ImageNet unsupervised data loader created')
return dataset, data_loader, dist_sampler
class ImageNet(torchvision.datasets.ImageFolder):
def __init__(
self,
root,
image_folder='imagenet_full_size/061417/',
tar_file='imagenet_full_size-061417.tar.gz',
transform=None,
train=True,
job_id=None,
local_rank=None,
copy_data=True,
index_targets=False
):
"""
ImageNet
Dataset wrapper (can copy data locally to machine)
:param root: root network directory for ImageNet data
:param image_folder: path to images inside root network directory
:param tar_file: zipped image_folder inside root network directory
:param train: whether to load train data (or validation)
:param job_id: scheduler job-id used to create dir on local machine
:param copy_data: whether to copy data from network file locally
:param index_targets: whether to index the id of each labeled image
"""
suffix = 'train/' if train else 'val/'
data_path = None
if copy_data:
logger.info('copying data locally')
data_path = copy_imgnt_locally(
root=root,
suffix=suffix,
image_folder=image_folder,
tar_file=tar_file,
job_id=job_id,
local_rank=local_rank)
if (not copy_data) or (data_path is None):
data_path = os.path.join(root, image_folder, suffix)
logger.info(f'data-path {data_path}')
super(ImageNet, self).__init__(root=data_path, transform=transform)
logger.info('Initialized ImageNet')
if index_targets:
self.targets = []
for sample in self.samples:
self.targets.append(sample[1])
self.targets = np.array(self.targets)
self.samples = np.array(self.samples)
mint = None
self.target_indices = []
for t in range(len(self.classes)):
indices = np.squeeze(np.argwhere(
self.targets == t)).tolist()
self.target_indices.append(indices)
mint = len(indices) if mint is None else min(mint, len(indices))
logger.debug(f'num-labeled target {t} {len(indices)}')
logger.info(f'min. labeled indices {mint}')
class ImageNetSubset(object):
def __init__(self, dataset, subset_file):
"""
ImageNetSubset
:param dataset: ImageNet dataset object
:param subset_file: '.txt' file containing IDs of IN1K images to keep
"""
self.dataset = dataset
self.subset_file = subset_file
self.filter_dataset_(subset_file)
def filter_dataset_(self, subset_file):
""" Filter self.dataset to a subset """
root = self.dataset.root
class_to_idx = self.dataset.class_to_idx
# -- update samples to subset of IN1k targets/samples
new_samples = []
logger.info(f'Using {subset_file}')
with open(subset_file, 'r') as rfile:
for line in rfile:
class_name = line.split('_')[0]
target = class_to_idx[class_name]
img = line.split('\n')[0]
new_samples.append(
(os.path.join(root, class_name, img), target)
)
self.samples = new_samples
@property
def classes(self):
return self.dataset.classes
def __len__(self):
return len(self.samples)
def __getitem__(self, index):
path, target = self.samples[index]
img = self.dataset.loader(path)
if self.dataset.transform is not None:
img = self.dataset.transform(img)
if self.dataset.target_transform is not None:
target = self.dataset.target_transform(target)
return img, target
def copy_imgnt_locally(
root,
suffix,
image_folder='imagenet_full_size/061417/',
tar_file='imagenet_full_size-061417.tar.gz',
job_id=None,
local_rank=None
):
if job_id is None:
try:
job_id = os.environ['SLURM_JOBID']
except Exception:
logger.info('No job-id, will load directly from network file')
return None
if local_rank is None:
try:
local_rank = int(os.environ['SLURM_LOCALID'])
except Exception:
logger.info('No job-id, will load directly from network file')
return None
source_file = os.path.join(root, tar_file)
target = f'/scratch/slurm_tmpdir/{job_id}/'
target_file = os.path.join(target, tar_file)
data_path = os.path.join(target, image_folder, suffix)
logger.info(f'{source_file}\n{target}\n{target_file}\n{data_path}')
tmp_sgnl_file = os.path.join(target, 'copy_signal.txt')
if not os.path.exists(data_path):
if local_rank == 0:
commands = [
['tar', '-xf', source_file, '-C', target]]
for cmnd in commands:
start_time = time.time()
logger.info(f'Executing {cmnd}')
subprocess.run(cmnd)
logger.info(f'Cmnd took {(time.time()-start_time)/60.} min.')
with open(tmp_sgnl_file, '+w') as f:
print('Done copying locally.', file=f)
else:
while not os.path.exists(tmp_sgnl_file):
time.sleep(60)
logger.info(f'{local_rank}: Checking {tmp_sgnl_file}')
return data_path
+156
View File
@@ -0,0 +1,156 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import logging
import sys
import torch
import src.models.vision_transformer as vit
from src.utils.schedulers import (
WarmupCosineSchedule,
CosineWDSchedule)
from src.utils.tensors import trunc_normal_
logging.basicConfig(stream=sys.stdout, level=logging.INFO)
logger = logging.getLogger()
def load_checkpoint(
device,
r_path,
encoder,
predictor,
target_encoder,
opt,
scaler,
):
try:
checkpoint = torch.load(r_path, map_location=torch.device('cpu'))
epoch = checkpoint['epoch']
# -- loading encoder
pretrained_dict = checkpoint['encoder']
msg = encoder.load_state_dict(pretrained_dict)
logger.info(f'loaded pretrained encoder from epoch {epoch} with msg: {msg}')
# -- loading predictor
pretrained_dict = checkpoint['predictor']
msg = predictor.load_state_dict(pretrained_dict)
logger.info(f'loaded pretrained encoder from epoch {epoch} with msg: {msg}')
# -- loading target_encoder
if target_encoder is not None:
print(list(checkpoint.keys()))
pretrained_dict = checkpoint['target_encoder']
msg = target_encoder.load_state_dict(pretrained_dict)
logger.info(f'loaded pretrained encoder from epoch {epoch} with msg: {msg}')
# -- loading optimizer
opt.load_state_dict(checkpoint['opt'])
if scaler is not None:
scaler.load_state_dict(checkpoint['scaler'])
logger.info(f'loaded optimizers from epoch {epoch}')
logger.info(f'read-path: {r_path}')
del checkpoint
except Exception as e:
logger.info(f'Encountered exception when loading checkpoint {e}')
epoch = 0
return encoder, predictor, target_encoder, opt, scaler, epoch
def init_model(
device,
patch_size=16,
model_name='vit_base',
crop_size=224,
pred_depth=6,
pred_emb_dim=384
):
encoder = vit.__dict__[model_name](
img_size=[crop_size],
patch_size=patch_size)
predictor = vit.__dict__['vit_predictor'](
num_patches=encoder.patch_embed.num_patches,
embed_dim=encoder.embed_dim,
predictor_embed_dim=pred_emb_dim,
depth=pred_depth,
num_heads=encoder.num_heads)
def init_weights(m):
if isinstance(m, torch.nn.Linear):
trunc_normal_(m.weight, std=0.02)
if m.bias is not None:
torch.nn.init.constant_(m.bias, 0)
elif isinstance(m, torch.nn.LayerNorm):
torch.nn.init.constant_(m.bias, 0)
torch.nn.init.constant_(m.weight, 1.0)
for m in encoder.modules():
init_weights(m)
for m in predictor.modules():
init_weights(m)
encoder.to(device)
predictor.to(device)
logger.info(encoder)
return encoder, predictor
def init_opt(
encoder,
predictor,
iterations_per_epoch,
start_lr,
ref_lr,
warmup,
num_epochs,
wd=1e-6,
final_wd=1e-6,
final_lr=0.0,
use_bfloat16=False,
ipe_scale=1.25
):
param_groups = [
{
'params': (p for n, p in encoder.named_parameters()
if ('bias' not in n) and (len(p.shape) != 1))
}, {
'params': (p for n, p in predictor.named_parameters()
if ('bias' not in n) and (len(p.shape) != 1))
}, {
'params': (p for n, p in encoder.named_parameters()
if ('bias' in n) or (len(p.shape) == 1)),
'WD_exclude': True,
'weight_decay': 0
}, {
'params': (p for n, p in predictor.named_parameters()
if ('bias' in n) or (len(p.shape) == 1)),
'WD_exclude': True,
'weight_decay': 0
}
]
logger.info('Using AdamW')
optimizer = torch.optim.AdamW(param_groups)
scheduler = WarmupCosineSchedule(
optimizer,
warmup_steps=int(warmup*iterations_per_epoch),
start_lr=start_lr,
ref_lr=ref_lr,
final_lr=final_lr,
T_max=int(ipe_scale*num_epochs*iterations_per_epoch))
wd_scheduler = CosineWDSchedule(
optimizer,
ref_wd=wd,
final_wd=final_wd,
T_max=int(ipe_scale*num_epochs*iterations_per_epoch))
scaler = torch.cuda.amp.GradScaler() if use_bfloat16 else None
return optimizer, scaler, scheduler, wd_scheduler
+21
View File
@@ -0,0 +1,21 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
from logging import getLogger
import torch
_GLOBAL_SEED = 0
logger = getLogger()
class DefaultCollator(object):
def __call__(self, batch):
collated_batch = torch.utils.data.default_collate(batch)
return collated_batch, None, None
+170
View File
@@ -0,0 +1,170 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import math
from multiprocessing import Value
from logging import getLogger
import torch
_GLOBAL_SEED = 0
logger = getLogger()
class MaskCollator(object):
def __init__(
self,
input_size=(224, 224),
patch_size=16,
enc_mask_scale=(0.2, 0.8),
pred_mask_scale=(0.2, 0.8),
aspect_ratio=(0.3, 3.0),
nenc=1,
npred=2,
min_keep=4,
allow_overlap=False
):
super(MaskCollator, self).__init__()
if not isinstance(input_size, tuple):
input_size = (input_size, ) * 2
self.patch_size = patch_size
self.height, self.width = input_size[0] // patch_size, input_size[1] // patch_size
self.enc_mask_scale = enc_mask_scale
self.pred_mask_scale = pred_mask_scale
self.aspect_ratio = aspect_ratio
self.nenc = nenc
self.npred = npred
self.min_keep = min_keep # minimum number of patches to keep
self.allow_overlap = allow_overlap # whether to allow overlap b/w enc and pred masks
self._itr_counter = Value('i', -1) # collator is shared across worker processes
def step(self):
i = self._itr_counter
with i.get_lock():
i.value += 1
v = i.value
return v
def _sample_block_size(self, generator, scale, aspect_ratio_scale):
_rand = torch.rand(1, generator=generator).item()
# -- Sample block scale
min_s, max_s = scale
mask_scale = min_s + _rand * (max_s - min_s)
max_keep = int(self.height * self.width * mask_scale)
# -- Sample block aspect-ratio
min_ar, max_ar = aspect_ratio_scale
aspect_ratio = min_ar + _rand * (max_ar - min_ar)
# -- Compute block height and width (given scale and aspect-ratio)
h = int(round(math.sqrt(max_keep * aspect_ratio)))
w = int(round(math.sqrt(max_keep / aspect_ratio)))
while h >= self.height:
h -= 1
while w >= self.width:
w -= 1
return (h, w)
def _sample_block_mask(self, b_size, acceptable_regions=None):
h, w = b_size
def constrain_mask(mask, tries=0):
""" Helper to restrict given mask to a set of acceptable regions """
N = max(int(len(acceptable_regions)-tries), 0)
for k in range(N):
mask *= acceptable_regions[k]
# --
# -- Loop to sample masks until we find a valid one
tries = 0
timeout = og_timeout = 20
valid_mask = False
while not valid_mask:
# -- Sample block top-left corner
top = torch.randint(0, self.height - h, (1,))
left = torch.randint(0, self.width - w, (1,))
mask = torch.zeros((self.height, self.width), dtype=torch.int32)
mask[top:top+h, left:left+w] = 1
# -- Constrain mask to a set of acceptable regions
if acceptable_regions is not None:
constrain_mask(mask, tries)
mask = torch.nonzero(mask.flatten())
# -- If mask too small try again
valid_mask = len(mask) > self.min_keep
if not valid_mask:
timeout -= 1
if timeout == 0:
tries += 1
timeout = og_timeout
logger.warning(f'Mask generator says: "Valid mask not found, decreasing acceptable-regions [{tries}]"')
mask = mask.squeeze()
# --
mask_complement = torch.ones((self.height, self.width), dtype=torch.int32)
mask_complement[top:top+h, left:left+w] = 0
# --
return mask, mask_complement
def __call__(self, batch):
'''
Create encoder and predictor masks when collating imgs into a batch
# 1. sample enc block (size + location) using seed
# 2. sample pred block (size) using seed
# 3. sample several enc block locations for each image (w/o seed)
# 4. sample several pred block locations for each image (w/o seed)
# 5. return enc mask and pred mask
'''
B = len(batch)
collated_batch = torch.utils.data.default_collate(batch)
seed = self.step()
g = torch.Generator()
g.manual_seed(seed)
p_size = self._sample_block_size(
generator=g,
scale=self.pred_mask_scale,
aspect_ratio_scale=self.aspect_ratio)
e_size = self._sample_block_size(
generator=g,
scale=self.enc_mask_scale,
aspect_ratio_scale=(1., 1.))
collated_masks_pred, collated_masks_enc = [], []
min_keep_pred = self.height * self.width
min_keep_enc = self.height * self.width
for _ in range(B):
masks_p, masks_C = [], []
for _ in range(self.npred):
mask, mask_C = self._sample_block_mask(p_size)
masks_p.append(mask)
masks_C.append(mask_C)
min_keep_pred = min(min_keep_pred, len(mask))
collated_masks_pred.append(masks_p)
acceptable_regions = masks_C
try:
if self.allow_overlap:
acceptable_regions= None
except Exception as e:
logger.warning(f'Encountered exception in mask-generator {e}')
masks_e = []
for _ in range(self.nenc):
mask, _ = self._sample_block_mask(e_size, acceptable_regions=acceptable_regions)
masks_e.append(mask)
min_keep_enc = min(min_keep_enc, len(mask))
collated_masks_enc.append(masks_e)
collated_masks_pred = [[cm[:min_keep_pred] for cm in cm_list] for cm_list in collated_masks_pred]
collated_masks_pred = torch.utils.data.default_collate(collated_masks_pred)
# --
collated_masks_enc = [[cm[:min_keep_enc] for cm in cm_list] for cm_list in collated_masks_enc]
collated_masks_enc = torch.utils.data.default_collate(collated_masks_enc)
return collated_batch, collated_masks_enc, collated_masks_pred
+72
View File
@@ -0,0 +1,72 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
from multiprocessing import Value
from logging import getLogger
import torch
_GLOBAL_SEED = 0
logger = getLogger()
class MaskCollator(object):
def __init__(
self,
ratio=(0.4, 0.6),
input_size=(224, 224),
patch_size=16,
):
super(MaskCollator, self).__init__()
if not isinstance(input_size, tuple):
input_size = (input_size, ) * 2
self.patch_size = patch_size
self.height, self.width = input_size[0] // patch_size, input_size[1] // patch_size
self.ratio = ratio
self._itr_counter = Value('i', -1) # collator is shared across worker processes
def step(self):
i = self._itr_counter
with i.get_lock():
i.value += 1
v = i.value
return v
def __call__(self, batch):
'''
Create encoder and predictor masks when collating imgs into a batch
# 1. sample enc block (size + location) using seed
# 2. sample pred block (size) using seed
# 3. sample several enc block locations for each image (w/o seed)
# 4. sample several pred block locations for each image (w/o seed)
# 5. return enc mask and pred mask
'''
B = len(batch)
collated_batch = torch.utils.data.default_collate(batch)
seed = self.step()
g = torch.Generator()
g.manual_seed(seed)
ratio = self.ratio
ratio = ratio[0] + torch.rand(1, generator=g).item() * (ratio[1] - ratio[0])
num_patches = self.height * self.width
num_keep = int(num_patches * (1. - ratio))
collated_masks_pred, collated_masks_enc = [], []
for _ in range(B):
m = torch.randperm(num_patches)
collated_masks_enc.append([m[:num_keep]])
collated_masks_pred.append([m[num_keep:]])
collated_masks_pred = torch.utils.data.default_collate(collated_masks_pred)
collated_masks_enc = torch.utils.data.default_collate(collated_masks_enc)
return collated_batch, collated_masks_enc, collated_masks_pred
+20
View File
@@ -0,0 +1,20 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import torch
def apply_masks(x, masks):
"""
:param x: tensor of shape [B (batch-size), N (num-patches), D (feature-dim)]
:param masks: list of tensors containing indices of patches in [N] to keep
"""
all_x = []
for m in masks:
mask_keep = m.unsqueeze(-1).repeat(1, 1, x.size(-1))
all_x += [torch.gather(x, dim=1, index=mask_keep)]
return torch.cat(all_x, dim=0)
+500
View File
@@ -0,0 +1,500 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import math
from functools import partial
import numpy as np
import torch
import torch.nn as nn
from src.utils.tensors import (
trunc_normal_,
repeat_interleave_batch
)
from src.masks.utils import apply_masks
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False):
"""
grid_size: int of the grid height and width
return:
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
"""
grid_h = np.arange(grid_size, dtype=float)
grid_w = np.arange(grid_size, dtype=float)
grid = np.meshgrid(grid_w, grid_h) # here w goes first
grid = np.stack(grid, axis=0)
grid = grid.reshape([2, 1, grid_size, grid_size])
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
if cls_token:
pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
return pos_embed
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
assert embed_dim % 2 == 0
# use half of dimensions to encode grid_h
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
return emb
def get_1d_sincos_pos_embed(embed_dim, grid_size, cls_token=False):
"""
grid_size: int of the grid length
return:
pos_embed: [grid_size, embed_dim] or [1+grid_size, embed_dim] (w/ or w/o cls_token)
"""
grid = np.arange(grid_size, dtype=float)
pos_embed = get_1d_sincos_pos_embed_from_grid(embed_dim, grid)
if cls_token:
pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
return pos_embed
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
"""
embed_dim: output dimension for each position
pos: a list of positions to be encoded: size (M,)
out: (M, D)
"""
assert embed_dim % 2 == 0
omega = np.arange(embed_dim // 2, dtype=float)
omega /= embed_dim / 2.
omega = 1. / 10000**omega # (D/2,)
pos = pos.reshape(-1) # (M,)
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
emb_sin = np.sin(out) # (M, D/2)
emb_cos = np.cos(out) # (M, D/2)
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
return emb
def drop_path(x, drop_prob: float = 0., training: bool = False):
if drop_prob == 0. or not training:
return x
keep_prob = 1 - drop_prob
shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype, device=x.device)
random_tensor.floor_() # binarize
output = x.div(keep_prob) * random_tensor
return output
class DropPath(nn.Module):
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
"""
def __init__(self, drop_prob=None):
super(DropPath, self).__init__()
self.drop_prob = drop_prob
def forward(self, x):
return drop_path(x, self.drop_prob, self.training)
class MLP(nn.Module):
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
super().__init__()
out_features = out_features or in_features
hidden_features = hidden_features or in_features
self.fc1 = nn.Linear(in_features, hidden_features)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features, out_features)
self.drop = nn.Dropout(drop)
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x
class Attention(nn.Module):
def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0.):
super().__init__()
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = qk_scale or head_dim ** -0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
def forward(self, x):
B, N, C = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2]
attn = (q @ k.transpose(-2, -1)) * self.scale
attn = attn.softmax(dim=-1)
attn = self.attn_drop(attn)
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
x = self.proj(x)
x = self.proj_drop(x)
return x, attn
class Block(nn.Module):
def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm):
super().__init__()
self.norm1 = norm_layer(dim)
self.attn = Attention(
dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
self.norm2 = norm_layer(dim)
mlp_hidden_dim = int(dim * mlp_ratio)
self.mlp = MLP(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
def forward(self, x, return_attention=False):
y, attn = self.attn(self.norm1(x))
if return_attention:
return attn
x = x + self.drop_path(y)
x = x + self.drop_path(self.mlp(self.norm2(x)))
return x
class PatchEmbed(nn.Module):
""" Image to Patch Embedding
"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
super().__init__()
num_patches = (img_size // patch_size) * (img_size // patch_size)
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = num_patches
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x):
B, C, H, W = x.shape
x = self.proj(x).flatten(2).transpose(1, 2)
return x
class ConvEmbed(nn.Module):
"""
3x3 Convolution stems for ViT following ViTC models
"""
def __init__(self, channels, strides, img_size=224, in_chans=3, batch_norm=True):
super().__init__()
# Build the stems
stem = []
channels = [in_chans] + channels
for i in range(len(channels) - 2):
stem += [nn.Conv2d(channels[i], channels[i+1], kernel_size=3,
stride=strides[i], padding=1, bias=(not batch_norm))]
if batch_norm:
stem += [nn.BatchNorm2d(channels[i+1])]
stem += [nn.ReLU(inplace=True)]
stem += [nn.Conv2d(channels[-2], channels[-1], kernel_size=1, stride=strides[-1])]
self.stem = nn.Sequential(*stem)
# Comptute the number of patches
stride_prod = int(np.prod(strides))
self.num_patches = (img_size[0] // stride_prod)**2
def forward(self, x):
p = self.stem(x)
return p.flatten(2).transpose(1, 2)
class VisionTransformerPredictor(nn.Module):
""" Vision Transformer """
def __init__(
self,
num_patches,
embed_dim=768,
predictor_embed_dim=384,
depth=6,
num_heads=12,
mlp_ratio=4.0,
qkv_bias=True,
qk_scale=None,
drop_rate=0.0,
attn_drop_rate=0.0,
drop_path_rate=0.0,
norm_layer=nn.LayerNorm,
init_std=0.02,
**kwargs
):
super().__init__()
self.predictor_embed = nn.Linear(embed_dim, predictor_embed_dim, bias=True)
self.mask_token = nn.Parameter(torch.zeros(1, 1, predictor_embed_dim))
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule
# --
self.predictor_pos_embed = nn.Parameter(torch.zeros(1, num_patches, predictor_embed_dim),
requires_grad=False)
predictor_pos_embed = get_2d_sincos_pos_embed(self.predictor_pos_embed.shape[-1],
int(num_patches**.5),
cls_token=False)
self.predictor_pos_embed.data.copy_(torch.from_numpy(predictor_pos_embed).float().unsqueeze(0))
# --
self.predictor_blocks = nn.ModuleList([
Block(
dim=predictor_embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer)
for i in range(depth)])
self.predictor_norm = norm_layer(predictor_embed_dim)
self.predictor_proj = nn.Linear(predictor_embed_dim, embed_dim, bias=True)
# ------
self.init_std = init_std
trunc_normal_(self.mask_token, std=self.init_std)
self.apply(self._init_weights)
self.fix_init_weight()
def fix_init_weight(self):
def rescale(param, layer_id):
param.div_(math.sqrt(2.0 * layer_id))
for layer_id, layer in enumerate(self.predictor_blocks):
rescale(layer.attn.proj.weight.data, layer_id + 1)
rescale(layer.mlp.fc2.weight.data, layer_id + 1)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=self.init_std)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
elif isinstance(m, nn.Conv2d):
trunc_normal_(m.weight, std=self.init_std)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def forward(self, x, masks_x, masks):
assert (masks is not None) and (masks_x is not None), 'Cannot run predictor without mask indices'
if not isinstance(masks_x, list):
masks_x = [masks_x]
if not isinstance(masks, list):
masks = [masks]
# -- Batch Size
B = len(x) // len(masks_x)
# -- map from encoder-dim to pedictor-dim
x = self.predictor_embed(x)
# -- add positional embedding to x tokens
x_pos_embed = self.predictor_pos_embed.repeat(B, 1, 1)
x += apply_masks(x_pos_embed, masks_x)
_, N_ctxt, D = x.shape
# -- concat mask tokens to x
pos_embs = self.predictor_pos_embed.repeat(B, 1, 1)
pos_embs = apply_masks(pos_embs, masks)
pos_embs = repeat_interleave_batch(pos_embs, B, repeat=len(masks_x))
# --
pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1)
# --
pred_tokens += pos_embs
x = x.repeat(len(masks), 1, 1)
x = torch.cat([x, pred_tokens], dim=1)
# -- fwd prop
for blk in self.predictor_blocks:
x = blk(x)
x = self.predictor_norm(x)
# -- return preds for mask tokens
x = x[:, N_ctxt:]
x = self.predictor_proj(x)
return x
class VisionTransformer(nn.Module):
""" Vision Transformer """
def __init__(
self,
img_size=[224],
patch_size=16,
in_chans=3,
embed_dim=768,
predictor_embed_dim=384,
depth=12,
predictor_depth=12,
num_heads=12,
mlp_ratio=4.0,
qkv_bias=True,
qk_scale=None,
drop_rate=0.0,
attn_drop_rate=0.0,
drop_path_rate=0.0,
norm_layer=nn.LayerNorm,
init_std=0.02,
**kwargs
):
super().__init__()
self.num_features = self.embed_dim = embed_dim
self.num_heads = num_heads
# --
self.patch_embed = PatchEmbed(
img_size=img_size[0],
patch_size=patch_size,
in_chans=in_chans,
embed_dim=embed_dim)
num_patches = self.patch_embed.num_patches
# --
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim), requires_grad=False)
pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1],
int(self.patch_embed.num_patches**.5),
cls_token=False)
self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
# --
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule
self.blocks = nn.ModuleList([
Block(
dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer)
for i in range(depth)])
self.norm = norm_layer(embed_dim)
# ------
self.init_std = init_std
self.apply(self._init_weights)
self.fix_init_weight()
def fix_init_weight(self):
def rescale(param, layer_id):
param.div_(math.sqrt(2.0 * layer_id))
for layer_id, layer in enumerate(self.blocks):
rescale(layer.attn.proj.weight.data, layer_id + 1)
rescale(layer.mlp.fc2.weight.data, layer_id + 1)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=self.init_std)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
elif isinstance(m, nn.Conv2d):
trunc_normal_(m.weight, std=self.init_std)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def forward(self, x, masks=None):
if masks is not None:
if not isinstance(masks, list):
masks = [masks]
# -- patchify x
x = self.patch_embed(x)
B, N, D = x.shape
# -- add positional embedding to x
pos_embed = self.interpolate_pos_encoding(x, self.pos_embed)
x = x + pos_embed
# -- mask x
if masks is not None:
x = apply_masks(x, masks)
# -- fwd prop
for i, blk in enumerate(self.blocks):
x = blk(x)
if self.norm is not None:
x = self.norm(x)
return x
def interpolate_pos_encoding(self, x, pos_embed):
npatch = x.shape[1] - 1
N = pos_embed.shape[1] - 1
if npatch == N:
return pos_embed
class_emb = pos_embed[:, 0]
pos_embed = pos_embed[:, 1:]
dim = x.shape[-1]
pos_embed = nn.functional.interpolate(
pos_embed.reshape(1, int(math.sqrt(N)), int(math.sqrt(N)), dim).permute(0, 3, 1, 2),
scale_factor=math.sqrt(npatch / N),
mode='bicubic',
)
pos_embed = pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)
return torch.cat((class_emb.unsqueeze(0), pos_embed), dim=1)
def vit_predictor(**kwargs):
model = VisionTransformerPredictor(
mlp_ratio=4, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6),
**kwargs)
return model
def vit_tiny(patch_size=16, **kwargs):
model = VisionTransformer(
patch_size=patch_size, embed_dim=192, depth=12, num_heads=3, mlp_ratio=4,
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)
return model
def vit_small(patch_size=16, **kwargs):
model = VisionTransformer(
patch_size=patch_size, embed_dim=384, depth=12, num_heads=6, mlp_ratio=4,
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)
return model
def vit_base(patch_size=16, **kwargs):
model = VisionTransformer(
patch_size=patch_size, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4,
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)
return model
def vit_large(patch_size=16, **kwargs):
model = VisionTransformer(
patch_size=patch_size, embed_dim=1024, depth=24, num_heads=16, mlp_ratio=4,
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)
return model
def vit_huge(patch_size=16, **kwargs):
model = VisionTransformer(
patch_size=patch_size, embed_dim=1280, depth=32, num_heads=16, mlp_ratio=4,
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)
return model
def vit_giant(patch_size=16, **kwargs):
model = VisionTransformer(
patch_size=patch_size, embed_dim=1408, depth=40, num_heads=16, mlp_ratio=48/11,
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), **kwargs)
return model
VIT_EMBED_DIMS = {
'vit_tiny': 192,
'vit_small': 384,
'vit_base': 768,
'vit_large': 1024,
'vit_huge': 1280,
'vit_giant': 1408,
}
+379
View File
@@ -0,0 +1,379 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import os
# -- FOR DISTRIBUTED TRAINING ENSURE ONLY 1 DEVICE VISIBLE PER PROCESS
try:
# -- WARNING: IF DOING DISTRIBUTED TRAINING ON A NON-SLURM CLUSTER, MAKE
# -- SURE TO UPDATE THIS TO GET LOCAL-RANK ON NODE, OR ENSURE
# -- THAT YOUR JOBS ARE LAUNCHED WITH ONLY 1 DEVICE VISIBLE
# -- TO EACH PROCESS
os.environ['CUDA_VISIBLE_DEVICES'] = os.environ['SLURM_LOCALID']
except Exception:
pass
import copy
import logging
import sys
import yaml
import numpy as np
import torch
import torch.multiprocessing as mp
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel
from src.masks.multiblock import MaskCollator as MBMaskCollator
from src.masks.utils import apply_masks
from src.utils.distributed import (
init_distributed,
AllReduce
)
from src.utils.logging import (
CSVLogger,
gpu_timer,
grad_logger,
AverageMeter)
from src.utils.tensors import repeat_interleave_batch
from src.datasets.imagenet1k import make_imagenet1k
from src.helper import (
load_checkpoint,
init_model,
init_opt)
from src.transforms import make_transforms
# --
log_timings = True
log_freq = 10
checkpoint_freq = 50
# --
_GLOBAL_SEED = 0
np.random.seed(_GLOBAL_SEED)
torch.manual_seed(_GLOBAL_SEED)
torch.backends.cudnn.benchmark = True
logging.basicConfig(stream=sys.stdout, level=logging.INFO)
logger = logging.getLogger()
def main(args, resume_preempt=False):
# ----------------------------------------------------------------------- #
# PASSED IN PARAMS FROM CONFIG FILE
# ----------------------------------------------------------------------- #
# -- META
use_bfloat16 = args['meta']['use_bfloat16']
model_name = args['meta']['model_name']
load_model = args['meta']['load_checkpoint'] or resume_preempt
r_file = args['meta']['read_checkpoint']
copy_data = args['meta']['copy_data']
pred_depth = args['meta']['pred_depth']
pred_emb_dim = args['meta']['pred_emb_dim']
if not torch.cuda.is_available():
device = torch.device('cpu')
else:
device = torch.device('cuda:0')
torch.cuda.set_device(device)
# -- DATA
use_gaussian_blur = args['data']['use_gaussian_blur']
use_horizontal_flip = args['data']['use_horizontal_flip']
use_color_distortion = args['data']['use_color_distortion']
color_jitter = args['data']['color_jitter_strength']
# --
batch_size = args['data']['batch_size']
pin_mem = args['data']['pin_mem']
num_workers = args['data']['num_workers']
root_path = args['data']['root_path']
image_folder = args['data']['image_folder']
crop_size = args['data']['crop_size']
crop_scale = args['data']['crop_scale']
# --
# -- MASK
allow_overlap = args['mask']['allow_overlap'] # whether to allow overlap b/w context and target blocks
patch_size = args['mask']['patch_size'] # patch-size for model training
num_enc_masks = args['mask']['num_enc_masks'] # number of context blocks
min_keep = args['mask']['min_keep'] # min number of patches in context block
enc_mask_scale = args['mask']['enc_mask_scale'] # scale of context blocks
num_pred_masks = args['mask']['num_pred_masks'] # number of target blocks
pred_mask_scale = args['mask']['pred_mask_scale'] # scale of target blocks
aspect_ratio = args['mask']['aspect_ratio'] # aspect ratio of target blocks
# --
# -- OPTIMIZATION
ema = args['optimization']['ema']
ipe_scale = args['optimization']['ipe_scale'] # scheduler scale factor (def: 1.0)
wd = float(args['optimization']['weight_decay'])
final_wd = float(args['optimization']['final_weight_decay'])
num_epochs = args['optimization']['epochs']
warmup = args['optimization']['warmup']
start_lr = args['optimization']['start_lr']
lr = args['optimization']['lr']
final_lr = args['optimization']['final_lr']
# -- LOGGING
folder = args['logging']['folder']
tag = args['logging']['write_tag']
dump = os.path.join(folder, 'params-ijepa.yaml')
with open(dump, 'w') as f:
yaml.dump(args, f)
# ----------------------------------------------------------------------- #
try:
mp.set_start_method('spawn')
except Exception:
pass
# -- init torch distributed backend
world_size, rank = init_distributed()
logger.info(f'Initialized (rank/world-size) {rank}/{world_size}')
if rank > 0:
logger.setLevel(logging.ERROR)
# -- log/checkpointing paths
log_file = os.path.join(folder, f'{tag}_r{rank}.csv')
save_path = os.path.join(folder, f'{tag}' + '-ep{epoch}.pth.tar')
latest_path = os.path.join(folder, f'{tag}-latest.pth.tar')
load_path = None
if load_model:
load_path = os.path.join(folder, r_file) if r_file is not None else latest_path
# -- make csv_logger
csv_logger = CSVLogger(log_file,
('%d', 'epoch'),
('%d', 'itr'),
('%.5f', 'loss'),
('%.5f', 'mask-A'),
('%.5f', 'mask-B'),
('%d', 'time (ms)'))
# -- init model
encoder, predictor = init_model(
device=device,
patch_size=patch_size,
crop_size=crop_size,
pred_depth=pred_depth,
pred_emb_dim=pred_emb_dim,
model_name=model_name)
target_encoder = copy.deepcopy(encoder)
# -- make data transforms
mask_collator = MBMaskCollator(
input_size=crop_size,
patch_size=patch_size,
pred_mask_scale=pred_mask_scale,
enc_mask_scale=enc_mask_scale,
aspect_ratio=aspect_ratio,
nenc=num_enc_masks,
npred=num_pred_masks,
allow_overlap=allow_overlap,
min_keep=min_keep)
transform = make_transforms(
crop_size=crop_size,
crop_scale=crop_scale,
gaussian_blur=use_gaussian_blur,
horizontal_flip=use_horizontal_flip,
color_distortion=use_color_distortion,
color_jitter=color_jitter)
# -- init data-loaders/samplers
_, unsupervised_loader, unsupervised_sampler = make_imagenet1k(
transform=transform,
batch_size=batch_size,
collator=mask_collator,
pin_mem=pin_mem,
training=True,
num_workers=num_workers,
world_size=world_size,
rank=rank,
root_path=root_path,
image_folder=image_folder,
copy_data=copy_data,
drop_last=True)
ipe = len(unsupervised_loader)
# -- init optimizer and scheduler
optimizer, scaler, scheduler, wd_scheduler = init_opt(
encoder=encoder,
predictor=predictor,
wd=wd,
final_wd=final_wd,
start_lr=start_lr,
ref_lr=lr,
final_lr=final_lr,
iterations_per_epoch=ipe,
warmup=warmup,
num_epochs=num_epochs,
ipe_scale=ipe_scale,
use_bfloat16=use_bfloat16)
encoder = DistributedDataParallel(encoder, static_graph=True)
predictor = DistributedDataParallel(predictor, static_graph=True)
target_encoder = DistributedDataParallel(target_encoder)
for p in target_encoder.parameters():
p.requires_grad = False
# -- momentum schedule
momentum_scheduler = (ema[0] + i*(ema[1]-ema[0])/(ipe*num_epochs*ipe_scale)
for i in range(int(ipe*num_epochs*ipe_scale)+1))
start_epoch = 0
# -- load training checkpoint
if load_model:
encoder, predictor, target_encoder, optimizer, scaler, start_epoch = load_checkpoint(
device=device,
r_path=load_path,
encoder=encoder,
predictor=predictor,
target_encoder=target_encoder,
opt=optimizer,
scaler=scaler)
for _ in range(start_epoch*ipe):
scheduler.step()
wd_scheduler.step()
next(momentum_scheduler)
mask_collator.step()
def save_checkpoint(epoch):
save_dict = {
'encoder': encoder.state_dict(),
'predictor': predictor.state_dict(),
'target_encoder': target_encoder.state_dict(),
'opt': optimizer.state_dict(),
'scaler': None if scaler is None else scaler.state_dict(),
'epoch': epoch,
'loss': loss_meter.avg,
'batch_size': batch_size,
'world_size': world_size,
'lr': lr
}
if rank == 0:
torch.save(save_dict, latest_path)
if (epoch + 1) % checkpoint_freq == 0:
torch.save(save_dict, save_path.format(epoch=f'{epoch + 1}'))
# -- TRAINING LOOP
for epoch in range(start_epoch, num_epochs):
logger.info('Epoch %d' % (epoch + 1))
# -- update distributed-data-loader epoch
unsupervised_sampler.set_epoch(epoch)
loss_meter = AverageMeter()
maskA_meter = AverageMeter()
maskB_meter = AverageMeter()
time_meter = AverageMeter()
for itr, (udata, masks_enc, masks_pred) in enumerate(unsupervised_loader):
def load_imgs():
# -- unsupervised imgs
imgs = udata[0].to(device, non_blocking=True)
masks_1 = [u.to(device, non_blocking=True) for u in masks_enc]
masks_2 = [u.to(device, non_blocking=True) for u in masks_pred]
return (imgs, masks_1, masks_2)
imgs, masks_enc, masks_pred = load_imgs()
maskA_meter.update(len(masks_enc[0][0]))
maskB_meter.update(len(masks_pred[0][0]))
def train_step():
_new_lr = scheduler.step()
_new_wd = wd_scheduler.step()
# --
def forward_target():
with torch.no_grad():
h = target_encoder(imgs)
h = F.layer_norm(h, (h.size(-1),)) # normalize over feature-dim
B = len(h)
# -- create targets (masked regions of h)
h = apply_masks(h, masks_pred)
h = repeat_interleave_batch(h, B, repeat=len(masks_enc))
return h
def forward_context():
z = encoder(imgs, masks_enc)
z = predictor(z, masks_enc, masks_pred)
return z
def loss_fn(z, h):
loss = F.smooth_l1_loss(z, h)
loss = AllReduce.apply(loss)
return loss
# Step 1. Forward
with torch.cuda.amp.autocast(dtype=torch.bfloat16, enabled=use_bfloat16):
h = forward_target()
z = forward_context()
loss = loss_fn(z, h)
# Step 2. Backward & step
if use_bfloat16:
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
else:
loss.backward()
optimizer.step()
grad_stats = grad_logger(encoder.named_parameters())
optimizer.zero_grad()
# Step 3. momentum update of target encoder
with torch.no_grad():
m = next(momentum_scheduler)
for param_q, param_k in zip(encoder.parameters(), target_encoder.parameters()):
param_k.data.mul_(m).add_((1.-m) * param_q.detach().data)
return (float(loss), _new_lr, _new_wd, grad_stats)
(loss, _new_lr, _new_wd, grad_stats), etime = gpu_timer(train_step)
loss_meter.update(loss)
time_meter.update(etime)
# -- Logging
def log_stats():
csv_logger.log(epoch + 1, itr, loss, maskA_meter.val, maskB_meter.val, etime)
if (itr % log_freq == 0) or np.isnan(loss) or np.isinf(loss):
logger.info('[%d, %5d] loss: %.3f '
'masks: %.1f %.1f '
'[wd: %.2e] [lr: %.2e] '
'[mem: %.2e] '
'(%.1f ms)'
% (epoch + 1, itr,
loss_meter.avg,
maskA_meter.avg,
maskB_meter.avg,
_new_wd,
_new_lr,
torch.cuda.max_memory_allocated() / 1024.**2,
time_meter.avg))
if grad_stats is not None:
logger.info('[%d, %5d] grad_stats: [%.2e %.2e] (%.2e, %.2e)'
% (epoch + 1, itr,
grad_stats.first_layer,
grad_stats.last_layer,
grad_stats.min,
grad_stats.max))
log_stats()
assert not np.isnan(loss), 'loss is nan'
# -- Save Checkpoint after every epoch
logger.info('avg. loss %.3f' % loss_meter.avg)
save_checkpoint(epoch+1)
if __name__ == "__main__":
main()
+67
View File
@@ -0,0 +1,67 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
from logging import getLogger
from PIL import ImageFilter
import torch
import torchvision.transforms as transforms
_GLOBAL_SEED = 0
logger = getLogger()
def make_transforms(
crop_size=224,
crop_scale=(0.3, 1.0),
color_jitter=1.0,
horizontal_flip=False,
color_distortion=False,
gaussian_blur=False,
normalization=((0.485, 0.456, 0.406),
(0.229, 0.224, 0.225))
):
logger.info('making imagenet data transforms')
def get_color_distortion(s=1.0):
# s is the strength of color distortion.
color_jitter = transforms.ColorJitter(0.8*s, 0.8*s, 0.8*s, 0.2*s)
rnd_color_jitter = transforms.RandomApply([color_jitter], p=0.8)
rnd_gray = transforms.RandomGrayscale(p=0.2)
color_distort = transforms.Compose([
rnd_color_jitter,
rnd_gray])
return color_distort
transform_list = []
transform_list += [transforms.RandomResizedCrop(crop_size, scale=crop_scale)]
if horizontal_flip:
transform_list += [transforms.RandomHorizontalFlip()]
if color_distortion:
transform_list += [get_color_distortion(s=color_jitter)]
if gaussian_blur:
transform_list += [GaussianBlur(p=0.5)]
transform_list += [transforms.ToTensor()]
transform_list += [transforms.Normalize(normalization[0], normalization[1])]
transform = transforms.Compose(transform_list)
return transform
class GaussianBlur(object):
def __init__(self, p=0.5, radius_min=0.1, radius_max=2.):
self.prob = p
self.radius_min = radius_min
self.radius_max = radius_max
def __call__(self, img):
if torch.bernoulli(torch.tensor(self.prob)) == 0:
return img
radius = self.radius_min + torch.rand(1) * (self.radius_max - self.radius_min)
return img.filter(ImageFilter.GaussianBlur(radius=radius))
+112
View File
@@ -0,0 +1,112 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import os
import torch
import torch.distributed as dist
from logging import getLogger
logger = getLogger()
def init_distributed(port=40112, rank_and_world_size=(None, None)):
if dist.is_available() and dist.is_initialized():
return dist.get_world_size(), dist.get_rank()
rank, world_size = rank_and_world_size
os.environ['MASTER_ADDR'] = 'localhost'
if (rank is None) or (world_size is None):
try:
world_size = int(os.environ['SLURM_NTASKS'])
rank = int(os.environ['SLURM_PROCID'])
os.environ['MASTER_ADDR'] = os.environ['HOSTNAME']
except Exception:
logger.info('SLURM vars not set (distributed training not available)')
world_size, rank = 1, 0
return world_size, rank
try:
os.environ['MASTER_PORT'] = str(port)
torch.distributed.init_process_group(
backend='nccl',
world_size=world_size,
rank=rank)
except Exception as e:
world_size, rank = 1, 0
logger.info(f'distributed training not available {e}')
return world_size, rank
class AllGather(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
if (
dist.is_available()
and dist.is_initialized()
and (dist.get_world_size() > 1)
):
x = x.contiguous()
outputs = [torch.zeros_like(x) for _ in range(dist.get_world_size())]
dist.all_gather(outputs, x)
return torch.cat(outputs, 0)
return x
@staticmethod
def backward(ctx, grads):
if (
dist.is_available()
and dist.is_initialized()
and (dist.get_world_size() > 1)
):
s = (grads.shape[0] // dist.get_world_size()) * dist.get_rank()
e = (grads.shape[0] // dist.get_world_size()) * (dist.get_rank() + 1)
grads = grads.contiguous()
dist.all_reduce(grads)
return grads[s:e]
return grads
class AllReduceSum(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
if (
dist.is_available()
and dist.is_initialized()
and (dist.get_world_size() > 1)
):
x = x.contiguous()
dist.all_reduce(x)
return x
@staticmethod
def backward(ctx, grads):
return grads
class AllReduce(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
if (
dist.is_available()
and dist.is_initialized()
and (dist.get_world_size() > 1)
):
x = x.contiguous() / dist.get_world_size()
dist.all_reduce(x)
return x
@staticmethod
def backward(ctx, grads):
return grads
+92
View File
@@ -0,0 +1,92 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import torch
def gpu_timer(closure, log_timings=True):
""" Helper to time gpu-time to execute closure() """
log_timings = log_timings and torch.cuda.is_available()
elapsed_time = -1.
if log_timings:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
result = closure()
if log_timings:
end.record()
torch.cuda.synchronize()
elapsed_time = start.elapsed_time(end)
return result, elapsed_time
class CSVLogger(object):
def __init__(self, fname, *argv):
self.fname = fname
self.types = []
# -- print headers
with open(self.fname, '+a') as f:
for i, v in enumerate(argv, 1):
self.types.append(v[0])
if i < len(argv):
print(v[1], end=',', file=f)
else:
print(v[1], end='\n', file=f)
def log(self, *argv):
with open(self.fname, '+a') as f:
for i, tv in enumerate(zip(self.types, argv), 1):
end = ',' if i < len(argv) else '\n'
print(tv[0] % tv[1], end=end, file=f)
class AverageMeter(object):
"""computes and stores the average and current value"""
def __init__(self):
self.reset()
def reset(self):
self.val = 0
self.avg = 0
self.max = float('-inf')
self.min = float('inf')
self.sum = 0
self.count = 0
def update(self, val, n=1):
self.val = val
try:
self.max = max(val, self.max)
self.min = min(val, self.min)
except Exception:
pass
self.sum += val * n
self.count += n
self.avg = self.sum / self.count
def grad_logger(named_params):
stats = AverageMeter()
stats.first_layer = None
stats.last_layer = None
for n, p in named_params:
if (p.grad is not None) and not (n.endswith('.bias') or len(p.shape) == 1):
grad_norm = float(torch.norm(p.grad.data))
stats.update(grad_norm)
if 'qkv' in n:
stats.last_layer = grad_norm
if stats.first_layer is None:
stats.first_layer = grad_norm
if stats.first_layer is None or stats.last_layer is None:
stats.first_layer = stats.last_layer = 0.
return stats
+76
View File
@@ -0,0 +1,76 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import math
class WarmupCosineSchedule(object):
def __init__(
self,
optimizer,
warmup_steps,
start_lr,
ref_lr,
T_max,
last_epoch=-1,
final_lr=0.
):
self.optimizer = optimizer
self.start_lr = start_lr
self.ref_lr = ref_lr
self.final_lr = final_lr
self.warmup_steps = warmup_steps
self.T_max = T_max - warmup_steps
self._step = 0.
def step(self):
self._step += 1
if self._step < self.warmup_steps:
progress = float(self._step) / float(max(1, self.warmup_steps))
new_lr = self.start_lr + progress * (self.ref_lr - self.start_lr)
else:
# -- progress after warmup
progress = float(self._step - self.warmup_steps) / float(max(1, self.T_max))
new_lr = max(self.final_lr,
self.final_lr + (self.ref_lr - self.final_lr) * 0.5 * (1. + math.cos(math.pi * progress)))
for group in self.optimizer.param_groups:
group['lr'] = new_lr
return new_lr
class CosineWDSchedule(object):
def __init__(
self,
optimizer,
ref_wd,
T_max,
final_wd=0.
):
self.optimizer = optimizer
self.ref_wd = ref_wd
self.final_wd = final_wd
self.T_max = T_max
self._step = 0.
def step(self):
self._step += 1
progress = self._step / self.T_max
new_wd = self.final_wd + (self.ref_wd - self.final_wd) * 0.5 * (1. + math.cos(math.pi * progress))
if self.final_wd <= self.ref_wd:
new_wd = max(self.final_wd, new_wd)
else:
new_wd = min(self.final_wd, new_wd)
for group in self.optimizer.param_groups:
if ('WD_exclude' not in group) or not group['WD_exclude']:
group['weight_decay'] = new_wd
return new_wd
+71
View File
@@ -0,0 +1,71 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
#
import math
import torch
from logging import getLogger
logger = getLogger()
def _no_grad_trunc_normal_(tensor, mean, std, a, b):
# Cut & paste from PyTorch official master until it's in a few official releases - RW
# Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf
def norm_cdf(x):
# Computes standard normal cumulative distribution function
return (1. + math.erf(x / math.sqrt(2.))) / 2.
with torch.no_grad():
# Values are generated by using a truncated uniform distribution and
# then using the inverse CDF for the normal distribution.
# Get upper and lower cdf values
l = norm_cdf((a - mean) / std)
u = norm_cdf((b - mean) / std)
# Uniformly fill tensor with values from [l, u], then translate to
# [2l-1, 2u-1].
tensor.uniform_(2 * l - 1, 2 * u - 1)
# Use inverse cdf transform for normal distribution to get truncated
# standard normal
tensor.erfinv_()
# Transform to proper mean, std
tensor.mul_(std * math.sqrt(2.))
tensor.add_(mean)
# Clamp to ensure it's in the proper range
tensor.clamp_(min=a, max=b)
return tensor
def trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.):
# type: (Tensor, float, float, float, float) -> Tensor
return _no_grad_trunc_normal_(tensor, mean, std, a, b)
def apply_masks(x, masks):
"""
:param x: tensor of shape [B (batch-size), N (num-patches), D (feature-dim)]
:param masks: list of tensors containing indices of patches in [N] to keep
"""
all_x = []
for m in masks:
mask_keep = m.unsqueeze(-1).repeat(1, 1, x.size(-1))
all_x += [torch.gather(x, dim=1, index=mask_keep)]
return torch.cat(all_x, dim=0)
def repeat_interleave_batch(x, B, repeat):
N = len(x) // B
x = torch.cat([
torch.cat([x[i*B:(i+1)*B] for _ in range(repeat)], dim=0)
for i in range(N)
], dim=0)
return x