From 45269d16bb858a7342465485dcc6a38133896854 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Thu, 4 Jun 2026 14:09:14 +0200 Subject: [PATCH] gradient checkpointing --- configs/wilds_vith16-448_ep300.yaml | 1 + src/helper.py | 9 ++++++--- src/models/vision_transformer.py | 15 +++++++++++++-- src/train.py | 4 +++- 4 files changed, 23 insertions(+), 6 deletions(-) diff --git a/configs/wilds_vith16-448_ep300.yaml b/configs/wilds_vith16-448_ep300.yaml index cfcefe6..9cc9e58 100644 --- a/configs/wilds_vith16-448_ep300.yaml +++ b/configs/wilds_vith16-448_ep300.yaml @@ -38,6 +38,7 @@ meta: pred_emb_dim: 384 read_checkpoint: null use_bfloat16: true + use_gradient_checkpointing: true optimization: ema: - 0.996 diff --git a/src/helper.py b/src/helper.py index dddf65d..70605aa 100644 --- a/src/helper.py +++ b/src/helper.py @@ -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): diff --git a/src/models/vision_transformer.py b/src/models/vision_transformer.py index 5fd49e0..d1a4222 100644 --- a/src/models/vision_transformer.py +++ b/src/models/vision_transformer.py @@ -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) diff --git a/src/train.py b/src/train.py index 12875f7..bfb5c35 100644 --- a/src/train.py +++ b/src/train.py @@ -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():