adapt params to ijepa paper

This commit is contained in:
YannAhlgrim
2026-06-02 12:33:51 +02:00
parent 628daab3b7
commit f2b501d8c9
6 changed files with 61 additions and 17 deletions
+5 -4
View File
@@ -1,5 +1,5 @@
data: data:
batch_size: 16 batch_size: 256
color_jitter_strength: 0.0 color_jitter_strength: 0.0
crop_scale: crop_scale:
- 1.0 - 1.0
@@ -13,7 +13,7 @@ data:
use_horizontal_flip: false use_horizontal_flip: false
use_random_resized_crop: false use_random_resized_crop: false
logging: logging:
folder: experiment_logs/vitb16.448-bs.16-ep.600 folder: experiment_logs/vitb16.448-bs.16-ep.600-paper
write_tag: jepa write_tag: jepa
mask: mask:
allow_overlap: false allow_overlap: false
@@ -48,6 +48,7 @@ optimization:
final_weight_decay: 0.4 final_weight_decay: 0.4
ipe_scale: 1.0 ipe_scale: 1.0
lr: 0.001 lr: 0.001
start_lr: 0.0002 start_lr: 0.0001
warmup: 40 warmup: 15
wd_schedule: linear
weight_decay: 0.04 weight_decay: 0.04
+4 -3
View File
@@ -1,5 +1,5 @@
data: data:
batch_size: 128 batch_size: 256
color_jitter_strength: 0.0 color_jitter_strength: 0.0
crop_scale: crop_scale:
- 0.3 - 0.3
@@ -46,6 +46,7 @@ optimization:
final_weight_decay: 0.4 final_weight_decay: 0.4
ipe_scale: 1.0 ipe_scale: 1.0
lr: 0.001 lr: 0.001
start_lr: 0.0002 start_lr: 0.0001
warmup: 40 warmup: 15
wd_schedule: linear
weight_decay: 0.04 weight_decay: 0.04
+4 -3
View File
@@ -1,5 +1,5 @@
data: data:
batch_size: 128 batch_size: 256
color_jitter_strength: 0.0 color_jitter_strength: 0.0
crop_scale: crop_scale:
- 1.0 - 1.0
@@ -47,6 +47,7 @@ optimization:
final_weight_decay: 0.4 final_weight_decay: 0.4
ipe_scale: 1.0 ipe_scale: 1.0
lr: 0.001 lr: 0.001
start_lr: 0.0002 start_lr: 0.0001
warmup: 40 warmup: 15
wd_schedule: linear
weight_decay: 0.04 weight_decay: 0.04
+19 -7
View File
@@ -13,7 +13,8 @@ import torch
import src.models.vision_transformer as vit import src.models.vision_transformer as vit
from src.utils.schedulers import ( from src.utils.schedulers import (
WarmupCosineSchedule, WarmupCosineSchedule,
CosineWDSchedule) CosineWDSchedule,
LinearWDSchedule)
from src.utils.tensors import trunc_normal_ from src.utils.tensors import trunc_normal_
logging.basicConfig(stream=sys.stdout, level=logging.INFO) logging.basicConfig(stream=sys.stdout, level=logging.INFO)
@@ -116,7 +117,8 @@ def init_opt(
final_wd=1e-6, final_wd=1e-6,
final_lr=0.0, final_lr=0.0,
use_bfloat16=False, use_bfloat16=False,
ipe_scale=1.25 ipe_scale=1.25,
wd_schedule='cosine'
): ):
param_groups = [ param_groups = [
{ {
@@ -147,10 +149,20 @@ def init_opt(
ref_lr=ref_lr, ref_lr=ref_lr,
final_lr=final_lr, final_lr=final_lr,
T_max=int(ipe_scale*num_epochs*iterations_per_epoch)) T_max=int(ipe_scale*num_epochs*iterations_per_epoch))
wd_scheduler = CosineWDSchedule( wd_schedule = str(wd_schedule).lower()
optimizer, if wd_schedule == 'linear':
ref_wd=wd, logger.info('Using linear weight-decay schedule')
final_wd=final_wd, wd_scheduler = LinearWDSchedule(
T_max=int(ipe_scale*num_epochs*iterations_per_epoch)) 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 scaler = torch.cuda.amp.GradScaler() if use_bfloat16 else None
return optimizer, scaler, scheduler, wd_scheduler return optimizer, scaler, scheduler, wd_scheduler
+2
View File
@@ -107,6 +107,7 @@ def main(args, resume_preempt=False):
ipe_scale = args["optimization"]["ipe_scale"] # scheduler scale factor (def: 1.0) ipe_scale = args["optimization"]["ipe_scale"] # scheduler scale factor (def: 1.0)
wd = float(args["optimization"]["weight_decay"]) wd = float(args["optimization"]["weight_decay"])
final_wd = float(args["optimization"]["final_weight_decay"]) final_wd = float(args["optimization"]["final_weight_decay"])
wd_schedule = args["optimization"].get("wd_schedule", "cosine")
num_epochs = args["optimization"]["epochs"] num_epochs = args["optimization"]["epochs"]
warmup = args["optimization"]["warmup"] warmup = args["optimization"]["warmup"]
start_lr = args["optimization"]["start_lr"] start_lr = args["optimization"]["start_lr"]
@@ -214,6 +215,7 @@ def main(args, resume_preempt=False):
num_epochs=num_epochs, num_epochs=num_epochs,
ipe_scale=ipe_scale, ipe_scale=ipe_scale,
use_bfloat16=use_bfloat16, use_bfloat16=use_bfloat16,
wd_schedule=wd_schedule,
) )
if dist.is_available() and dist.is_initialized() and world_size > 1: if dist.is_available() and dist.is_initialized() and world_size > 1:
encoder = DistributedDataParallel(encoder, static_graph=True) encoder = DistributedDataParallel(encoder, static_graph=True)
+27
View File
@@ -74,3 +74,30 @@ class CosineWDSchedule(object):
if ('WD_exclude' not in group) or not group['WD_exclude']: if ('WD_exclude' not in group) or not group['WD_exclude']:
group['weight_decay'] = new_wd group['weight_decay'] = new_wd
return 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