adapt params to ijepa paper
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+14
-2
@@ -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,6 +149,16 @@ 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_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(
|
wd_scheduler = CosineWDSchedule(
|
||||||
optimizer,
|
optimizer,
|
||||||
ref_wd=wd,
|
ref_wd=wd,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user