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