add lars opt and configs
This commit is contained in:
@@ -8,7 +8,7 @@ meta:
|
|||||||
num_classes: 182
|
num_classes: 182
|
||||||
|
|
||||||
data:
|
data:
|
||||||
batch_size: 128
|
batch_size: 16384
|
||||||
root_path: ./wilds_data
|
root_path: ./wilds_data
|
||||||
num_workers: 10
|
num_workers: 10
|
||||||
pin_mem: true
|
pin_mem: true
|
||||||
@@ -24,26 +24,33 @@ mask:
|
|||||||
patch_size: 14
|
patch_size: 14
|
||||||
|
|
||||||
optimization:
|
optimization:
|
||||||
optimizer: adamw # 'adamw' or 'sgd'
|
optimizer: lars # 'adamw', 'sgd', or 'lars'
|
||||||
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: 300 # can be set higher if early_stopping
|
epochs: 50
|
||||||
lr: 5.0e-4
|
lr: 0.01
|
||||||
weight_decay: 1.0e-2
|
weight_decay: 5.0e-4
|
||||||
use_cosine_schedule: true
|
use_cosine_schedule: false
|
||||||
start_lr: 0.0001
|
lr_schedule: step
|
||||||
final_lr: 1.0e-06
|
step_milestones: [15, 30, 45]
|
||||||
warmup: 5
|
step_gamma: 0.1
|
||||||
|
start_lr: 0.0
|
||||||
|
final_lr: 0.0
|
||||||
|
warmup: 0
|
||||||
|
momentum: 0.9
|
||||||
|
lars_eta: 0.001
|
||||||
|
lars_eps: 1.0e-8
|
||||||
|
lars_exclude_bias_and_norm: true
|
||||||
ipe_scale: 1.0
|
ipe_scale: 1.0
|
||||||
early_stopping:
|
early_stopping:
|
||||||
enabled: true
|
enabled: false
|
||||||
patience: 10
|
patience: 10
|
||||||
min_delta: 1.0e-4
|
min_delta: 1.0e-4
|
||||||
min_epochs: 15
|
min_epochs: 50
|
||||||
restore_best_weights: true
|
restore_best_weights: true
|
||||||
|
|
||||||
validation:
|
validation:
|
||||||
eval_every: 1
|
eval_every: 1
|
||||||
|
|
||||||
logging:
|
logging:
|
||||||
folder: experiment_logs/supervised-vith14.224-bs.128-ep.300-lp-0crop/
|
folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.01-wd0.0005/
|
||||||
write_tag: linear_probe
|
write_tag: linear_probe-lars-lr0.01-wd0.0005
|
||||||
|
|||||||
@@ -0,0 +1,56 @@
|
|||||||
|
meta:
|
||||||
|
model_name: vit_huge
|
||||||
|
embed_dim: 1280
|
||||||
|
load_checkpoint: true
|
||||||
|
checkpoint_folder: experiment_logs/vith14.224-bs.128-ep.300/
|
||||||
|
read_checkpoint: jepa-ep300.pth.tar
|
||||||
|
use_bfloat16: true
|
||||||
|
num_classes: 182
|
||||||
|
|
||||||
|
data:
|
||||||
|
batch_size: 16384
|
||||||
|
root_path: ./wilds_data
|
||||||
|
num_workers: 10
|
||||||
|
pin_mem: true
|
||||||
|
crop_size: 224
|
||||||
|
crop_scale: [1.0, 1.0]
|
||||||
|
use_random_resized_crop: false
|
||||||
|
use_horizontal_flip: false
|
||||||
|
use_color_distortion: false
|
||||||
|
color_jitter_strength: 0.0
|
||||||
|
use_gaussian_blur: false
|
||||||
|
|
||||||
|
mask:
|
||||||
|
patch_size: 14
|
||||||
|
|
||||||
|
optimization:
|
||||||
|
optimizer: lars # 'adamw', 'sgd', or 'lars'
|
||||||
|
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||||
|
epochs: 50
|
||||||
|
lr: 0.001
|
||||||
|
weight_decay: 0.0
|
||||||
|
use_cosine_schedule: false
|
||||||
|
lr_schedule: step
|
||||||
|
step_milestones: [15, 30, 45]
|
||||||
|
step_gamma: 0.1
|
||||||
|
start_lr: 0.0
|
||||||
|
final_lr: 0.0
|
||||||
|
warmup: 0
|
||||||
|
momentum: 0.9
|
||||||
|
lars_eta: 0.001
|
||||||
|
lars_eps: 1.0e-8
|
||||||
|
lars_exclude_bias_and_norm: true
|
||||||
|
ipe_scale: 1.0
|
||||||
|
early_stopping:
|
||||||
|
enabled: false
|
||||||
|
patience: 10
|
||||||
|
min_delta: 1.0e-4
|
||||||
|
min_epochs: 50
|
||||||
|
restore_best_weights: true
|
||||||
|
|
||||||
|
validation:
|
||||||
|
eval_every: 1
|
||||||
|
|
||||||
|
logging:
|
||||||
|
folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.001-wd0.0/
|
||||||
|
write_tag: linear_probe-lars-lr0.001-wd0.0
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
meta:
|
||||||
|
model_name: vit_huge
|
||||||
|
embed_dim: 1280
|
||||||
|
load_checkpoint: true
|
||||||
|
checkpoint_folder: experiment_logs/vith14.224-bs.128-ep.300/
|
||||||
|
read_checkpoint: jepa-ep300.pth.tar
|
||||||
|
use_bfloat16: true
|
||||||
|
num_classes: 182
|
||||||
|
|
||||||
|
data:
|
||||||
|
batch_size: 16384
|
||||||
|
root_path: ./wilds_data
|
||||||
|
num_workers: 10
|
||||||
|
pin_mem: true
|
||||||
|
crop_size: 224
|
||||||
|
crop_scale: [1.0, 1.0]
|
||||||
|
use_random_resized_crop: false
|
||||||
|
use_horizontal_flip: false
|
||||||
|
use_color_distortion: false
|
||||||
|
color_jitter_strength: 0.0
|
||||||
|
use_gaussian_blur: false
|
||||||
|
|
||||||
|
mask:
|
||||||
|
patch_size: 14
|
||||||
|
|
||||||
|
optimization:
|
||||||
|
optimizer: lars # 'adamw', 'sgd', or 'lars'
|
||||||
|
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||||
|
epochs: 50
|
||||||
|
lr: 0.001
|
||||||
|
weight_decay: 5.0e-4
|
||||||
|
use_cosine_schedule: false
|
||||||
|
lr_schedule: step
|
||||||
|
step_milestones: [15, 30, 45]
|
||||||
|
step_gamma: 0.1
|
||||||
|
start_lr: 0.0
|
||||||
|
final_lr: 0.0
|
||||||
|
warmup: 0
|
||||||
|
momentum: 0.9
|
||||||
|
lars_eta: 0.001
|
||||||
|
lars_eps: 1.0e-8
|
||||||
|
lars_exclude_bias_and_norm: true
|
||||||
|
ipe_scale: 1.0
|
||||||
|
early_stopping:
|
||||||
|
enabled: false
|
||||||
|
patience: 10
|
||||||
|
min_delta: 1.0e-4
|
||||||
|
min_epochs: 50
|
||||||
|
restore_best_weights: true
|
||||||
|
|
||||||
|
validation:
|
||||||
|
eval_every: 1
|
||||||
|
|
||||||
|
logging:
|
||||||
|
folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.001-wd0.0005/
|
||||||
|
write_tag: linear_probe-lars-lr0.001-wd0.0005
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
meta:
|
||||||
|
model_name: vit_huge
|
||||||
|
embed_dim: 1280
|
||||||
|
load_checkpoint: true
|
||||||
|
checkpoint_folder: experiment_logs/vith14.224-bs.128-ep.300/
|
||||||
|
read_checkpoint: jepa-ep300.pth.tar
|
||||||
|
use_bfloat16: true
|
||||||
|
num_classes: 182
|
||||||
|
|
||||||
|
data:
|
||||||
|
batch_size: 16384
|
||||||
|
root_path: ./wilds_data
|
||||||
|
num_workers: 10
|
||||||
|
pin_mem: true
|
||||||
|
crop_size: 224
|
||||||
|
crop_scale: [1.0, 1.0]
|
||||||
|
use_random_resized_crop: false
|
||||||
|
use_horizontal_flip: false
|
||||||
|
use_color_distortion: false
|
||||||
|
color_jitter_strength: 0.0
|
||||||
|
use_gaussian_blur: false
|
||||||
|
|
||||||
|
mask:
|
||||||
|
patch_size: 14
|
||||||
|
|
||||||
|
optimization:
|
||||||
|
optimizer: lars # 'adamw', 'sgd', or 'lars'
|
||||||
|
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||||
|
epochs: 50
|
||||||
|
lr: 0.01
|
||||||
|
weight_decay: 0.0
|
||||||
|
use_cosine_schedule: false
|
||||||
|
lr_schedule: step
|
||||||
|
step_milestones: [15, 30, 45]
|
||||||
|
step_gamma: 0.1
|
||||||
|
start_lr: 0.0
|
||||||
|
final_lr: 0.0
|
||||||
|
warmup: 0
|
||||||
|
momentum: 0.9
|
||||||
|
lars_eta: 0.001
|
||||||
|
lars_eps: 1.0e-8
|
||||||
|
lars_exclude_bias_and_norm: true
|
||||||
|
ipe_scale: 1.0
|
||||||
|
early_stopping:
|
||||||
|
enabled: false
|
||||||
|
patience: 10
|
||||||
|
min_delta: 1.0e-4
|
||||||
|
min_epochs: 50
|
||||||
|
restore_best_weights: true
|
||||||
|
|
||||||
|
validation:
|
||||||
|
eval_every: 1
|
||||||
|
|
||||||
|
logging:
|
||||||
|
folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.01-wd0.0/
|
||||||
|
write_tag: linear_probe-lars-lr0.01-wd0.0
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
meta:
|
||||||
|
model_name: vit_huge
|
||||||
|
embed_dim: 1280
|
||||||
|
load_checkpoint: true
|
||||||
|
checkpoint_folder: experiment_logs/vith14.224-bs.128-ep.300/
|
||||||
|
read_checkpoint: jepa-ep300.pth.tar
|
||||||
|
use_bfloat16: true
|
||||||
|
num_classes: 182
|
||||||
|
|
||||||
|
data:
|
||||||
|
batch_size: 16384
|
||||||
|
root_path: ./wilds_data
|
||||||
|
num_workers: 10
|
||||||
|
pin_mem: true
|
||||||
|
crop_size: 224
|
||||||
|
crop_scale: [1.0, 1.0]
|
||||||
|
use_random_resized_crop: false
|
||||||
|
use_horizontal_flip: false
|
||||||
|
use_color_distortion: false
|
||||||
|
color_jitter_strength: 0.0
|
||||||
|
use_gaussian_blur: false
|
||||||
|
|
||||||
|
mask:
|
||||||
|
patch_size: 14
|
||||||
|
|
||||||
|
optimization:
|
||||||
|
optimizer: lars # 'adamw', 'sgd', or 'lars'
|
||||||
|
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||||
|
epochs: 50
|
||||||
|
lr: 0.05
|
||||||
|
weight_decay: 0.0
|
||||||
|
use_cosine_schedule: false
|
||||||
|
lr_schedule: step
|
||||||
|
step_milestones: [15, 30, 45]
|
||||||
|
step_gamma: 0.1
|
||||||
|
start_lr: 0.0
|
||||||
|
final_lr: 0.0
|
||||||
|
warmup: 0
|
||||||
|
momentum: 0.9
|
||||||
|
lars_eta: 0.001
|
||||||
|
lars_eps: 1.0e-8
|
||||||
|
lars_exclude_bias_and_norm: true
|
||||||
|
ipe_scale: 1.0
|
||||||
|
early_stopping:
|
||||||
|
enabled: false
|
||||||
|
patience: 10
|
||||||
|
min_delta: 1.0e-4
|
||||||
|
min_epochs: 50
|
||||||
|
restore_best_weights: true
|
||||||
|
|
||||||
|
validation:
|
||||||
|
eval_every: 1
|
||||||
|
|
||||||
|
logging:
|
||||||
|
folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.05-wd0.0/
|
||||||
|
write_tag: linear_probe-lars-lr0.05-wd0.0
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
meta:
|
||||||
|
model_name: vit_huge
|
||||||
|
embed_dim: 1280
|
||||||
|
load_checkpoint: true
|
||||||
|
checkpoint_folder: experiment_logs/vith14.224-bs.128-ep.300/
|
||||||
|
read_checkpoint: jepa-ep300.pth.tar
|
||||||
|
use_bfloat16: true
|
||||||
|
num_classes: 182
|
||||||
|
|
||||||
|
data:
|
||||||
|
batch_size: 16384
|
||||||
|
root_path: ./wilds_data
|
||||||
|
num_workers: 10
|
||||||
|
pin_mem: true
|
||||||
|
crop_size: 224
|
||||||
|
crop_scale: [1.0, 1.0]
|
||||||
|
use_random_resized_crop: false
|
||||||
|
use_horizontal_flip: false
|
||||||
|
use_color_distortion: false
|
||||||
|
color_jitter_strength: 0.0
|
||||||
|
use_gaussian_blur: false
|
||||||
|
|
||||||
|
mask:
|
||||||
|
patch_size: 14
|
||||||
|
|
||||||
|
optimization:
|
||||||
|
optimizer: lars # 'adamw', 'sgd', or 'lars'
|
||||||
|
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||||
|
epochs: 50
|
||||||
|
lr: 0.05
|
||||||
|
weight_decay: 5.0e-4
|
||||||
|
use_cosine_schedule: false
|
||||||
|
lr_schedule: step
|
||||||
|
step_milestones: [15, 30, 45]
|
||||||
|
step_gamma: 0.1
|
||||||
|
start_lr: 0.0
|
||||||
|
final_lr: 0.0
|
||||||
|
warmup: 0
|
||||||
|
momentum: 0.9
|
||||||
|
lars_eta: 0.001
|
||||||
|
lars_eps: 1.0e-8
|
||||||
|
lars_exclude_bias_and_norm: true
|
||||||
|
ipe_scale: 1.0
|
||||||
|
early_stopping:
|
||||||
|
enabled: false
|
||||||
|
patience: 10
|
||||||
|
min_delta: 1.0e-4
|
||||||
|
min_epochs: 50
|
||||||
|
restore_best_weights: true
|
||||||
|
|
||||||
|
validation:
|
||||||
|
eval_every: 1
|
||||||
|
|
||||||
|
logging:
|
||||||
|
folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.05-wd0.0005/
|
||||||
|
write_tag: linear_probe-lars-lr0.05-wd0.0005
|
||||||
+25
-3
@@ -16,6 +16,7 @@ from src.models.head import ViTClassifier
|
|||||||
from src.transforms import make_transforms, make_transform_eval
|
from src.transforms import make_transforms, make_transform_eval
|
||||||
from src.utils.distributed import init_distributed
|
from src.utils.distributed import init_distributed
|
||||||
from src.utils.logging import CSVLogger, AverageMeter
|
from src.utils.logging import CSVLogger, AverageMeter
|
||||||
|
from src.utils.optimizers import LARS
|
||||||
|
|
||||||
# --
|
# --
|
||||||
log_freq = 10
|
log_freq = 10
|
||||||
@@ -217,17 +218,38 @@ def main(args, resume_preempt=False):
|
|||||||
logger.info("Training full model (Fine-tuning mode)")
|
logger.info("Training full model (Fine-tuning mode)")
|
||||||
|
|
||||||
params = [p for p in model.parameters() if p.requires_grad]
|
params = [p for p in model.parameters() if p.requires_grad]
|
||||||
if o_args["optimizer"].lower() == "adamw":
|
optimizer_name = o_args["optimizer"].lower()
|
||||||
|
if optimizer_name == "adamw":
|
||||||
optimizer = torch.optim.AdamW(
|
optimizer = torch.optim.AdamW(
|
||||||
params, lr=o_args["lr"], weight_decay=o_args["weight_decay"]
|
params, lr=o_args["lr"], weight_decay=o_args["weight_decay"]
|
||||||
)
|
)
|
||||||
|
elif optimizer_name == "lars":
|
||||||
|
optimizer = LARS(
|
||||||
|
params,
|
||||||
|
lr=o_args["lr"],
|
||||||
|
weight_decay=o_args["weight_decay"],
|
||||||
|
momentum=o_args.get("momentum", 0.9),
|
||||||
|
eta=o_args.get("lars_eta", 0.001),
|
||||||
|
eps=o_args.get("lars_eps", 1e-8),
|
||||||
|
exclude_bias_and_norm=o_args.get("lars_exclude_bias_and_norm", True),
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
optimizer = torch.optim.SGD(
|
optimizer = torch.optim.SGD(
|
||||||
params, lr=o_args["lr"], momentum=0.9, weight_decay=o_args["weight_decay"]
|
params,
|
||||||
|
lr=o_args["lr"],
|
||||||
|
momentum=o_args.get("momentum", 0.9),
|
||||||
|
weight_decay=o_args["weight_decay"],
|
||||||
)
|
)
|
||||||
|
|
||||||
scheduler = None
|
scheduler = None
|
||||||
if o_args["use_cosine_schedule"]:
|
lr_schedule = o_args.get("lr_schedule", "cosine").lower()
|
||||||
|
if lr_schedule == "step":
|
||||||
|
scheduler = torch.optim.lr_scheduler.MultiStepLR(
|
||||||
|
optimizer,
|
||||||
|
milestones=o_args.get("step_milestones", [15, 30, 45]),
|
||||||
|
gamma=o_args.get("step_gamma", 0.1),
|
||||||
|
)
|
||||||
|
elif 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["final_lr"]
|
optimizer, T_max=o_args["epochs"], eta_min=o_args["final_lr"]
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,84 @@
|
|||||||
|
import math
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
class LARS(torch.optim.Optimizer):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr,
|
||||||
|
weight_decay=0.0,
|
||||||
|
momentum=0.9,
|
||||||
|
eta=0.001,
|
||||||
|
eps=1e-8,
|
||||||
|
exclude_bias_and_norm=True,
|
||||||
|
):
|
||||||
|
if lr <= 0.0:
|
||||||
|
raise ValueError(f"Invalid lr: {lr}")
|
||||||
|
if weight_decay < 0.0:
|
||||||
|
raise ValueError(f"Invalid weight_decay: {weight_decay}")
|
||||||
|
if momentum < 0.0:
|
||||||
|
raise ValueError(f"Invalid momentum: {momentum}")
|
||||||
|
if eta <= 0.0:
|
||||||
|
raise ValueError(f"Invalid eta: {eta}")
|
||||||
|
if eps <= 0.0:
|
||||||
|
raise ValueError(f"Invalid eps: {eps}")
|
||||||
|
|
||||||
|
defaults = dict(
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
eta=eta,
|
||||||
|
eps=eps,
|
||||||
|
exclude_bias_and_norm=exclude_bias_and_norm,
|
||||||
|
)
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
momentum = group["momentum"]
|
||||||
|
eta = group["eta"]
|
||||||
|
eps = group["eps"]
|
||||||
|
exclude_bias_and_norm = group["exclude_bias_and_norm"]
|
||||||
|
|
||||||
|
for p in group["params"]:
|
||||||
|
if p.grad is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
grad = p.grad
|
||||||
|
if grad.is_sparse:
|
||||||
|
raise RuntimeError("LARS does not support sparse gradients")
|
||||||
|
|
||||||
|
param_norm = torch.norm(p)
|
||||||
|
grad_norm = torch.norm(grad)
|
||||||
|
|
||||||
|
lars_lr = 1.0
|
||||||
|
if not exclude_bias_and_norm or p.ndim > 1:
|
||||||
|
if param_norm > 0.0 and grad_norm > 0.0:
|
||||||
|
lars_lr = eta * param_norm / (grad_norm + weight_decay * param_norm + eps)
|
||||||
|
|
||||||
|
d_p = grad
|
||||||
|
if weight_decay != 0.0 and (not exclude_bias_and_norm or p.ndim > 1):
|
||||||
|
d_p = d_p.add(p, alpha=weight_decay)
|
||||||
|
|
||||||
|
if momentum != 0.0:
|
||||||
|
param_state = self.state.setdefault(p, {})
|
||||||
|
if "momentum_buffer" not in param_state:
|
||||||
|
buf = param_state["momentum_buffer"] = torch.clone(d_p).detach()
|
||||||
|
else:
|
||||||
|
buf = param_state["momentum_buffer"]
|
||||||
|
buf.mul_(momentum).add_(d_p)
|
||||||
|
d_p = buf
|
||||||
|
|
||||||
|
p.add_(d_p, alpha=-lr * lars_lr)
|
||||||
|
|
||||||
|
return loss
|
||||||
Reference in New Issue
Block a user