adapt params to ijepa paper
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+19
-7
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user