gradient checkpointing
This commit is contained in:
@@ -38,6 +38,7 @@ meta:
|
|||||||
pred_emb_dim: 384
|
pred_emb_dim: 384
|
||||||
read_checkpoint: null
|
read_checkpoint: null
|
||||||
use_bfloat16: true
|
use_bfloat16: true
|
||||||
|
use_gradient_checkpointing: true
|
||||||
optimization:
|
optimization:
|
||||||
ema:
|
ema:
|
||||||
- 0.996
|
- 0.996
|
||||||
|
|||||||
+6
-3
@@ -72,17 +72,20 @@ def init_model(
|
|||||||
model_name='vit_base',
|
model_name='vit_base',
|
||||||
crop_size=224,
|
crop_size=224,
|
||||||
pred_depth=6,
|
pred_depth=6,
|
||||||
pred_emb_dim=384
|
pred_emb_dim=384,
|
||||||
|
use_gradient_checkpointing=False
|
||||||
):
|
):
|
||||||
encoder = vit.__dict__[model_name](
|
encoder = vit.__dict__[model_name](
|
||||||
img_size=[crop_size],
|
img_size=[crop_size],
|
||||||
patch_size=patch_size)
|
patch_size=patch_size,
|
||||||
|
use_checkpoint=use_gradient_checkpointing)
|
||||||
predictor = vit.__dict__['vit_predictor'](
|
predictor = vit.__dict__['vit_predictor'](
|
||||||
num_patches=encoder.patch_embed.num_patches,
|
num_patches=encoder.patch_embed.num_patches,
|
||||||
embed_dim=encoder.embed_dim,
|
embed_dim=encoder.embed_dim,
|
||||||
predictor_embed_dim=pred_emb_dim,
|
predictor_embed_dim=pred_emb_dim,
|
||||||
depth=pred_depth,
|
depth=pred_depth,
|
||||||
num_heads=encoder.num_heads)
|
num_heads=encoder.num_heads,
|
||||||
|
use_checkpoint=use_gradient_checkpointing)
|
||||||
|
|
||||||
def init_weights(m):
|
def init_weights(m):
|
||||||
if isinstance(m, torch.nn.Linear):
|
if isinstance(m, torch.nn.Linear):
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import numpy as np
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
from torch.utils.checkpoint import checkpoint as _checkpoint
|
||||||
|
|
||||||
from src.utils.tensors import (
|
from src.utils.tensors import (
|
||||||
trunc_normal_,
|
trunc_normal_,
|
||||||
@@ -82,6 +83,12 @@ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
|||||||
return emb
|
return emb
|
||||||
|
|
||||||
|
|
||||||
|
def _maybe_checkpoint(blk, x, use_checkpoint, training):
|
||||||
|
if use_checkpoint and training and x.requires_grad:
|
||||||
|
return _checkpoint(blk, x, use_reentrant=False)
|
||||||
|
return blk(x)
|
||||||
|
|
||||||
|
|
||||||
def drop_path(x, drop_prob: float = 0., training: bool = False):
|
def drop_path(x, drop_prob: float = 0., training: bool = False):
|
||||||
if drop_prob == 0. or not training:
|
if drop_prob == 0. or not training:
|
||||||
return x
|
return x
|
||||||
@@ -234,9 +241,11 @@ class VisionTransformerPredictor(nn.Module):
|
|||||||
drop_path_rate=0.0,
|
drop_path_rate=0.0,
|
||||||
norm_layer=nn.LayerNorm,
|
norm_layer=nn.LayerNorm,
|
||||||
init_std=0.02,
|
init_std=0.02,
|
||||||
|
use_checkpoint=False,
|
||||||
**kwargs
|
**kwargs
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
self.use_checkpoint = use_checkpoint
|
||||||
self.predictor_embed = nn.Linear(embed_dim, predictor_embed_dim, bias=True)
|
self.predictor_embed = nn.Linear(embed_dim, predictor_embed_dim, bias=True)
|
||||||
self.mask_token = nn.Parameter(torch.zeros(1, 1, predictor_embed_dim))
|
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
|
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule
|
||||||
@@ -316,7 +325,7 @@ class VisionTransformerPredictor(nn.Module):
|
|||||||
|
|
||||||
# -- fwd prop
|
# -- fwd prop
|
||||||
for blk in self.predictor_blocks:
|
for blk in self.predictor_blocks:
|
||||||
x = blk(x)
|
x = _maybe_checkpoint(blk, x, self.use_checkpoint, self.training)
|
||||||
x = self.predictor_norm(x)
|
x = self.predictor_norm(x)
|
||||||
|
|
||||||
# -- return preds for mask tokens
|
# -- return preds for mask tokens
|
||||||
@@ -346,11 +355,13 @@ class VisionTransformer(nn.Module):
|
|||||||
drop_path_rate=0.0,
|
drop_path_rate=0.0,
|
||||||
norm_layer=nn.LayerNorm,
|
norm_layer=nn.LayerNorm,
|
||||||
init_std=0.02,
|
init_std=0.02,
|
||||||
|
use_checkpoint=False,
|
||||||
**kwargs
|
**kwargs
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.num_features = self.embed_dim = embed_dim
|
self.num_features = self.embed_dim = embed_dim
|
||||||
self.num_heads = num_heads
|
self.num_heads = num_heads
|
||||||
|
self.use_checkpoint = use_checkpoint
|
||||||
# --
|
# --
|
||||||
self.patch_embed = PatchEmbed(
|
self.patch_embed = PatchEmbed(
|
||||||
img_size=img_size[0],
|
img_size=img_size[0],
|
||||||
@@ -424,7 +435,7 @@ class VisionTransformer(nn.Module):
|
|||||||
# -- fwd prop
|
# -- fwd prop
|
||||||
layer_outputs = []
|
layer_outputs = []
|
||||||
for i, blk in enumerate(self.blocks):
|
for i, blk in enumerate(self.blocks):
|
||||||
x = blk(x)
|
x = _maybe_checkpoint(blk, x, self.use_checkpoint, self.training)
|
||||||
if return_layer_outputs:
|
if return_layer_outputs:
|
||||||
layer_outputs.append(x)
|
layer_outputs.append(x)
|
||||||
|
|
||||||
|
|||||||
+3
-1
@@ -63,6 +63,7 @@ def main(args, resume_preempt=False):
|
|||||||
# -- META
|
# -- META
|
||||||
use_bfloat16 = args["meta"]["use_bfloat16"]
|
use_bfloat16 = args["meta"]["use_bfloat16"]
|
||||||
model_name = args["meta"]["model_name"]
|
model_name = args["meta"]["model_name"]
|
||||||
|
use_gradient_checkpointing = args["meta"].get("use_gradient_checkpointing", False)
|
||||||
load_model = args["meta"]["load_checkpoint"] or resume_preempt
|
load_model = args["meta"]["load_checkpoint"] or resume_preempt
|
||||||
r_file = args["meta"]["read_checkpoint"]
|
r_file = args["meta"]["read_checkpoint"]
|
||||||
copy_data = args["meta"]["copy_data"]
|
copy_data = args["meta"]["copy_data"]
|
||||||
@@ -161,6 +162,7 @@ def main(args, resume_preempt=False):
|
|||||||
pred_depth=pred_depth,
|
pred_depth=pred_depth,
|
||||||
pred_emb_dim=pred_emb_dim,
|
pred_emb_dim=pred_emb_dim,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
|
use_gradient_checkpointing=use_gradient_checkpointing,
|
||||||
)
|
)
|
||||||
target_encoder = copy.deepcopy(encoder)
|
target_encoder = copy.deepcopy(encoder)
|
||||||
|
|
||||||
@@ -339,7 +341,7 @@ def main(args, resume_preempt=False):
|
|||||||
loss.backward()
|
loss.backward()
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
grad_stats = grad_logger(encoder.named_parameters())
|
grad_stats = grad_logger(encoder.named_parameters())
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
# Step 3. momentum update of target encoder
|
# Step 3. momentum update of target encoder
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
|
|||||||
Reference in New Issue
Block a user