add lars opt and configs

This commit is contained in:
YannAhlgrim
2026-05-20 15:05:18 +02:00
parent feb8e225f2
commit 5ce42b8ad0
8 changed files with 409 additions and 16 deletions
+20 -13
View File
@@ -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
View File
@@ -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"]
) )
+84
View File
@@ -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