From f2b501d8c94b1517aaedacc24d28b61e84910cc1 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Tue, 2 Jun 2026 12:33:51 +0200 Subject: [PATCH] adapt params to ijepa paper --- configs/wilds_vitb16-448_ep600.yaml | 9 +++++---- configs/wilds_vith14_ep300.yaml | 7 ++++--- configs/wilds_vith16-448_ep300.yaml | 7 ++++--- src/helper.py | 26 +++++++++++++++++++------- src/train.py | 2 ++ src/utils/schedulers.py | 27 +++++++++++++++++++++++++++ 6 files changed, 61 insertions(+), 17 deletions(-) diff --git a/configs/wilds_vitb16-448_ep600.yaml b/configs/wilds_vitb16-448_ep600.yaml index 822ba23..38d013a 100644 --- a/configs/wilds_vitb16-448_ep600.yaml +++ b/configs/wilds_vitb16-448_ep600.yaml @@ -1,5 +1,5 @@ data: - batch_size: 16 + batch_size: 256 color_jitter_strength: 0.0 crop_scale: - 1.0 @@ -13,7 +13,7 @@ data: use_horizontal_flip: false use_random_resized_crop: false logging: - folder: experiment_logs/vitb16.448-bs.16-ep.600 + folder: experiment_logs/vitb16.448-bs.16-ep.600-paper write_tag: jepa mask: allow_overlap: false @@ -48,6 +48,7 @@ optimization: final_weight_decay: 0.4 ipe_scale: 1.0 lr: 0.001 - start_lr: 0.0002 - warmup: 40 + start_lr: 0.0001 + warmup: 15 + wd_schedule: linear weight_decay: 0.04 diff --git a/configs/wilds_vith14_ep300.yaml b/configs/wilds_vith14_ep300.yaml index bfba219..6016e8b 100644 --- a/configs/wilds_vith14_ep300.yaml +++ b/configs/wilds_vith14_ep300.yaml @@ -1,5 +1,5 @@ data: - batch_size: 128 + batch_size: 256 color_jitter_strength: 0.0 crop_scale: - 0.3 @@ -46,6 +46,7 @@ optimization: final_weight_decay: 0.4 ipe_scale: 1.0 lr: 0.001 - start_lr: 0.0002 - warmup: 40 + start_lr: 0.0001 + warmup: 15 + wd_schedule: linear weight_decay: 0.04 diff --git a/configs/wilds_vith16-448_ep300.yaml b/configs/wilds_vith16-448_ep300.yaml index 1f8332e..cfcefe6 100644 --- a/configs/wilds_vith16-448_ep300.yaml +++ b/configs/wilds_vith16-448_ep300.yaml @@ -1,5 +1,5 @@ data: - batch_size: 128 + batch_size: 256 color_jitter_strength: 0.0 crop_scale: - 1.0 @@ -47,6 +47,7 @@ optimization: final_weight_decay: 0.4 ipe_scale: 1.0 lr: 0.001 - start_lr: 0.0002 - warmup: 40 + start_lr: 0.0001 + warmup: 15 + wd_schedule: linear weight_decay: 0.04 diff --git a/src/helper.py b/src/helper.py index 5af8cd3..dddf65d 100644 --- a/src/helper.py +++ b/src/helper.py @@ -13,7 +13,8 @@ import torch import src.models.vision_transformer as vit from src.utils.schedulers import ( WarmupCosineSchedule, - CosineWDSchedule) + CosineWDSchedule, + LinearWDSchedule) from src.utils.tensors import trunc_normal_ logging.basicConfig(stream=sys.stdout, level=logging.INFO) @@ -116,7 +117,8 @@ def init_opt( final_wd=1e-6, final_lr=0.0, use_bfloat16=False, - ipe_scale=1.25 + ipe_scale=1.25, + wd_schedule='cosine' ): param_groups = [ { @@ -147,10 +149,20 @@ def init_opt( 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)) + wd_schedule = str(wd_schedule).lower() + if wd_schedule == 'linear': + logger.info('Using linear weight-decay schedule') + wd_scheduler = LinearWDSchedule( + optimizer, + ref_wd=wd, + final_wd=final_wd, + T_max=int(ipe_scale*num_epochs*iterations_per_epoch)) + else: + logger.info('Using cosine weight-decay schedule') + 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 diff --git a/src/train.py b/src/train.py index 413742e..12875f7 100644 --- a/src/train.py +++ b/src/train.py @@ -107,6 +107,7 @@ def main(args, resume_preempt=False): 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"]) + wd_schedule = args["optimization"].get("wd_schedule", "cosine") num_epochs = args["optimization"]["epochs"] warmup = args["optimization"]["warmup"] start_lr = args["optimization"]["start_lr"] @@ -214,6 +215,7 @@ def main(args, resume_preempt=False): num_epochs=num_epochs, ipe_scale=ipe_scale, use_bfloat16=use_bfloat16, + wd_schedule=wd_schedule, ) if dist.is_available() and dist.is_initialized() and world_size > 1: encoder = DistributedDataParallel(encoder, static_graph=True) diff --git a/src/utils/schedulers.py b/src/utils/schedulers.py index df02e2b..2ee446f 100644 --- a/src/utils/schedulers.py +++ b/src/utils/schedulers.py @@ -74,3 +74,30 @@ class CosineWDSchedule(object): if ('WD_exclude' not in group) or not group['WD_exclude']: group['weight_decay'] = new_wd return new_wd + + +class LinearWDSchedule(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 + progress = min(max(progress, 0.0), 1.0) + new_wd = self.ref_wd + progress * (self.final_wd - self.ref_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