diff --git a/configs/supervised_wilds_vith14_ep300-lp.yaml b/configs/supervised_wilds_vith14_ep300-lp.yaml index 61168e7..9bfeb19 100644 --- a/configs/supervised_wilds_vith14_ep300-lp.yaml +++ b/configs/supervised_wilds_vith14_ep300-lp.yaml @@ -8,7 +8,7 @@ meta: num_classes: 182 data: - batch_size: 128 + batch_size: 16384 root_path: ./wilds_data num_workers: 10 pin_mem: true @@ -24,26 +24,33 @@ mask: patch_size: 14 optimization: - optimizer: adamw # 'adamw' or 'sgd' + optimizer: lars # 'adamw', 'sgd', or 'lars' freeze_weights: true # true for linear probing, false for full fine-tuning - epochs: 300 # can be set higher if early_stopping - lr: 5.0e-4 - weight_decay: 1.0e-2 - use_cosine_schedule: true - start_lr: 0.0001 - final_lr: 1.0e-06 - warmup: 5 + epochs: 50 + lr: 0.01 + 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: true + enabled: false patience: 10 min_delta: 1.0e-4 - min_epochs: 15 + min_epochs: 50 restore_best_weights: true validation: eval_every: 1 logging: - folder: experiment_logs/supervised-vith14.224-bs.128-ep.300-lp-0crop/ - write_tag: linear_probe + folder: experiment_logs/supervised-vith14.224-bs.16384-ep.50-lp-lars-lr0.01-wd0.0005/ + write_tag: linear_probe-lars-lr0.01-wd0.0005 diff --git a/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.001-wd0.0.yaml b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.001-wd0.0.yaml new file mode 100644 index 0000000..5e5225a --- /dev/null +++ b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.001-wd0.0.yaml @@ -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 diff --git a/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.001-wd0.0005.yaml b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.001-wd0.0005.yaml new file mode 100644 index 0000000..c9c9807 --- /dev/null +++ b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.001-wd0.0005.yaml @@ -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 diff --git a/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.01-wd0.0.yaml b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.01-wd0.0.yaml new file mode 100644 index 0000000..dd00493 --- /dev/null +++ b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.01-wd0.0.yaml @@ -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 diff --git a/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.05-wd0.0.yaml b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.05-wd0.0.yaml new file mode 100644 index 0000000..2d9b08e --- /dev/null +++ b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.05-wd0.0.yaml @@ -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 diff --git a/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.05-wd0.0005.yaml b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.05-wd0.0005.yaml new file mode 100644 index 0000000..c9f9879 --- /dev/null +++ b/configs/supervised_wilds_vith14_ep50-lp-lars-lr0.05-wd0.0005.yaml @@ -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 diff --git a/src/train_supervised.py b/src/train_supervised.py index 8b49286..8950952 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -16,6 +16,7 @@ from src.models.head import ViTClassifier from src.transforms import make_transforms, make_transform_eval from src.utils.distributed import init_distributed from src.utils.logging import CSVLogger, AverageMeter +from src.utils.optimizers import LARS # -- log_freq = 10 @@ -217,17 +218,38 @@ def main(args, resume_preempt=False): logger.info("Training full model (Fine-tuning mode)") 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( 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: 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 - 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( optimizer, T_max=o_args["epochs"], eta_min=o_args["final_lr"] ) diff --git a/src/utils/optimizers.py b/src/utils/optimizers.py new file mode 100644 index 0000000..df9c9b6 --- /dev/null +++ b/src/utils/optimizers.py @@ -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