add checkpoint_folder + tweak params for supervised fine tuning + remove default values of params
This commit is contained in:
@@ -2,7 +2,8 @@ meta:
|
|||||||
model_name: vit_huge
|
model_name: vit_huge
|
||||||
embed_dim: 1280
|
embed_dim: 1280
|
||||||
load_checkpoint: true
|
load_checkpoint: true
|
||||||
read_checkpoint: jepa-latest.pth.tar
|
checkpoint_folder: experiment_logs/wilds-vith14.224-bs.128-ep.300/
|
||||||
|
read_checkpoint: jepa-ep300.pth.tar
|
||||||
use_bfloat16: true
|
use_bfloat16: true
|
||||||
num_classes: 182
|
num_classes: 182
|
||||||
|
|
||||||
@@ -12,23 +13,31 @@ data:
|
|||||||
num_workers: 10
|
num_workers: 10
|
||||||
pin_mem: true
|
pin_mem: true
|
||||||
crop_size: 224
|
crop_size: 224
|
||||||
|
crop_scale: [0.3, 1.0]
|
||||||
|
use_horizontal_flip: false
|
||||||
|
use_color_distortion: false
|
||||||
|
color_jitter_strength: 0.0
|
||||||
|
use_gaussian_blur: false
|
||||||
|
|
||||||
|
mask:
|
||||||
|
patch_size: 14
|
||||||
|
|
||||||
optimization:
|
optimization:
|
||||||
optimizer: adamw # 'adamw' or 'sgd'
|
optimizer: adamw # 'adamw' or 'sgd'
|
||||||
freeze_weights: true # true for linear probing, false for full fine-tuning
|
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||||
epochs: 200 # can be set higher if early_stopping
|
epochs: 300 # can be set higher if early_stopping
|
||||||
lr: 0.001
|
lr: 5.0e-4
|
||||||
weight_decay: 0.05
|
weight_decay: 1.0e-2
|
||||||
use_cosine_schedule: false
|
use_cosine_schedule: true
|
||||||
start_lr: 0.0001
|
start_lr: 0.0001
|
||||||
final_lr: 1.0e-06
|
final_lr: 1.0e-06
|
||||||
warmup: 5
|
warmup: 5
|
||||||
ipe_scale: 1.0
|
ipe_scale: 1.0
|
||||||
early_stopping:
|
early_stopping:
|
||||||
enabled: true
|
enabled: true
|
||||||
patience: 6
|
patience: 10
|
||||||
min_delta: 0.0
|
min_delta: 1.0e-4
|
||||||
min_epochs: 5
|
min_epochs: 15
|
||||||
restore_best_weights: true
|
restore_best_weights: true
|
||||||
|
|
||||||
validation:
|
validation:
|
||||||
|
|||||||
+32
-25
@@ -169,9 +169,10 @@ def main(args, resume_preempt=False):
|
|||||||
m_args = args["meta"]
|
m_args = args["meta"]
|
||||||
o_args = args["optimization"]
|
o_args = args["optimization"]
|
||||||
d_args = args["data"]
|
d_args = args["data"]
|
||||||
|
mk_args = args["mask"]
|
||||||
l_args = args["logging"]
|
l_args = args["logging"]
|
||||||
v_args = args.get("validation", {})
|
v_args = args["validation"]
|
||||||
es_args = o_args.get("early_stopping", {})
|
es_args = o_args["early_stopping"]
|
||||||
|
|
||||||
folder = l_args["folder"]
|
folder = l_args["folder"]
|
||||||
tag = l_args["write_tag"]
|
tag = l_args["write_tag"]
|
||||||
@@ -199,12 +200,12 @@ def main(args, resume_preempt=False):
|
|||||||
|
|
||||||
encoder, _ = init_model(
|
encoder, _ = init_model(
|
||||||
device=device,
|
device=device,
|
||||||
patch_size=args.get("mask", {}).get("patch_size", 14),
|
patch_size=mk_args["patch_size"],
|
||||||
crop_size=d_args["crop_size"],
|
crop_size=d_args["crop_size"],
|
||||||
model_name=m_args["model_name"],
|
model_name=m_args["model_name"],
|
||||||
)
|
)
|
||||||
|
|
||||||
embed_dim = m_args.get("embed_dim")
|
embed_dim = m_args["embed_dim"]
|
||||||
model = ViTClassifier(encoder, m_args["num_classes"], embed_dim).to(device)
|
model = ViTClassifier(encoder, m_args["num_classes"], embed_dim).to(device)
|
||||||
|
|
||||||
if o_args["freeze_weights"]:
|
if o_args["freeze_weights"]:
|
||||||
@@ -226,18 +227,18 @@ def main(args, resume_preempt=False):
|
|||||||
)
|
)
|
||||||
|
|
||||||
scheduler = None
|
scheduler = None
|
||||||
if o_args.get("use_cosine_schedule", False):
|
if o_args["use_cosine_schedule"]:
|
||||||
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||||
optimizer, T_max=o_args["epochs"], eta_min=o_args.get("final_lr", 0.0)
|
optimizer, T_max=o_args["epochs"], eta_min=o_args["final_lr"]
|
||||||
)
|
)
|
||||||
|
|
||||||
train_transform = make_transforms(
|
train_transform = make_transforms(
|
||||||
crop_size=d_args["crop_size"],
|
crop_size=d_args["crop_size"],
|
||||||
crop_scale=tuple(d_args.get("crop_scale", (0.3, 1.0))),
|
crop_scale=tuple(d_args["crop_scale"]),
|
||||||
horizontal_flip=d_args.get("use_horizontal_flip", False),
|
horizontal_flip=d_args["use_horizontal_flip"],
|
||||||
color_distortion=d_args.get("use_color_distortion", False),
|
color_distortion=d_args["use_color_distortion"],
|
||||||
color_jitter=d_args.get("color_jitter_strength", 1.0),
|
color_jitter=d_args["color_jitter_strength"],
|
||||||
gaussian_blur=d_args.get("use_gaussian_blur", False),
|
gaussian_blur=d_args["use_gaussian_blur"],
|
||||||
)
|
)
|
||||||
val_transform = make_transform_eval(
|
val_transform = make_transform_eval(
|
||||||
crop_size=d_args["crop_size"],
|
crop_size=d_args["crop_size"],
|
||||||
@@ -251,8 +252,8 @@ def main(args, resume_preempt=False):
|
|||||||
rank=rank,
|
rank=rank,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
collator=None,
|
collator=None,
|
||||||
num_workers=d_args.get("num_workers", 8),
|
num_workers=d_args["num_workers"],
|
||||||
pin_mem=d_args.get("pin_mem", True),
|
pin_mem=d_args["pin_mem"],
|
||||||
drop_last=True,
|
drop_last=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -264,8 +265,8 @@ def main(args, resume_preempt=False):
|
|||||||
rank=rank,
|
rank=rank,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
collator=None,
|
collator=None,
|
||||||
num_workers=d_args.get("num_workers", 8),
|
num_workers=d_args["num_workers"],
|
||||||
pin_mem=d_args.get("pin_mem", True),
|
pin_mem=d_args["pin_mem"],
|
||||||
drop_last=False,
|
drop_last=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -273,11 +274,11 @@ def main(args, resume_preempt=False):
|
|||||||
model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()])
|
model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()])
|
||||||
|
|
||||||
early_stopper = EarlyStopping(
|
early_stopper = EarlyStopping(
|
||||||
enabled=es_args.get("enabled", False),
|
enabled=es_args["enabled"],
|
||||||
patience=es_args.get("patience", 10),
|
patience=es_args["patience"],
|
||||||
min_delta=es_args.get("min_delta", 0.0),
|
min_delta=es_args["min_delta"],
|
||||||
min_epochs=es_args.get("min_epochs", 0),
|
min_epochs=es_args["min_epochs"],
|
||||||
restore_best_weights=es_args.get("restore_best_weights", True),
|
restore_best_weights=es_args["restore_best_weights"],
|
||||||
)
|
)
|
||||||
|
|
||||||
start_epoch = 0
|
start_epoch = 0
|
||||||
@@ -287,10 +288,16 @@ def main(args, resume_preempt=False):
|
|||||||
if os.path.exists(latest_path):
|
if os.path.exists(latest_path):
|
||||||
checkpoint_to_load = latest_path
|
checkpoint_to_load = latest_path
|
||||||
resuming_interrupted = True
|
resuming_interrupted = True
|
||||||
elif m_args.get("load_checkpoint", False):
|
elif m_args["load_checkpoint"]:
|
||||||
r_file = m_args.get("read_checkpoint")
|
r_file = m_args["read_checkpoint"]
|
||||||
if r_file is not None:
|
checkpoint_folder = m_args["checkpoint_folder"]
|
||||||
checkpoint_to_load = os.path.join(folder, r_file)
|
if os.path.isabs(r_file):
|
||||||
|
checkpoint_to_load = r_file
|
||||||
|
else:
|
||||||
|
checkpoint_to_load = os.path.join(checkpoint_folder, r_file)
|
||||||
|
|
||||||
|
if checkpoint_to_load is not None and not os.path.exists(checkpoint_to_load):
|
||||||
|
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_to_load}")
|
||||||
|
|
||||||
if checkpoint_to_load and os.path.exists(checkpoint_to_load):
|
if checkpoint_to_load and os.path.exists(checkpoint_to_load):
|
||||||
checkpoint = torch.load(checkpoint_to_load, map_location="cpu")
|
checkpoint = torch.load(checkpoint_to_load, map_location="cpu")
|
||||||
@@ -342,7 +349,7 @@ def main(args, resume_preempt=False):
|
|||||||
if is_best:
|
if is_best:
|
||||||
torch.save(save_dict, best_path)
|
torch.save(save_dict, best_path)
|
||||||
|
|
||||||
eval_every = int(v_args.get("eval_every", 1))
|
eval_every = int(v_args["eval_every"])
|
||||||
|
|
||||||
for epoch in range(start_epoch, o_args["epochs"]):
|
for epoch in range(start_epoch, o_args["epochs"]):
|
||||||
train_sampler.set_epoch(epoch)
|
train_sampler.set_epoch(epoch)
|
||||||
|
|||||||
Reference in New Issue
Block a user