gradient checkpointing

This commit is contained in:
YannAhlgrim
2026-06-04 14:09:14 +02:00
parent f2b501d8c9
commit 45269d16bb
4 changed files with 23 additions and 6 deletions
+6 -3
View File
@@ -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):
+13 -2
View File
@@ -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
View File
@@ -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():