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
+19 -7
View File
@@ -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
+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)
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)
+27
View File
@@ -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