From 75cf127d74c17cd2bda4e17968222e4906c85765 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Mon, 27 Apr 2026 11:14:14 +0200 Subject: [PATCH] add checkpoint_folder + tweak params for supervised fine tuning + remove default values of params --- configs/supervised_wilds_vith14_ep300.yaml | 25 +++++++--- src/train_supervised.py | 57 ++++++++++++---------- 2 files changed, 49 insertions(+), 33 deletions(-) diff --git a/configs/supervised_wilds_vith14_ep300.yaml b/configs/supervised_wilds_vith14_ep300.yaml index 4319394..675491c 100644 --- a/configs/supervised_wilds_vith14_ep300.yaml +++ b/configs/supervised_wilds_vith14_ep300.yaml @@ -2,7 +2,8 @@ meta: model_name: vit_huge embed_dim: 1280 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 num_classes: 182 @@ -12,23 +13,31 @@ data: num_workers: 10 pin_mem: true 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: optimizer: adamw # 'adamw' or 'sgd' freeze_weights: true # true for linear probing, false for full fine-tuning - epochs: 200 # can be set higher if early_stopping - lr: 0.001 - weight_decay: 0.05 - use_cosine_schedule: false + 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 ipe_scale: 1.0 early_stopping: enabled: true - patience: 6 - min_delta: 0.0 - min_epochs: 5 + patience: 10 + min_delta: 1.0e-4 + min_epochs: 15 restore_best_weights: true validation: diff --git a/src/train_supervised.py b/src/train_supervised.py index 05517f5..94c3b58 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -169,9 +169,10 @@ def main(args, resume_preempt=False): m_args = args["meta"] o_args = args["optimization"] d_args = args["data"] + mk_args = args["mask"] l_args = args["logging"] - v_args = args.get("validation", {}) - es_args = o_args.get("early_stopping", {}) + v_args = args["validation"] + es_args = o_args["early_stopping"] folder = l_args["folder"] tag = l_args["write_tag"] @@ -199,12 +200,12 @@ def main(args, resume_preempt=False): encoder, _ = init_model( device=device, - patch_size=args.get("mask", {}).get("patch_size", 14), + patch_size=mk_args["patch_size"], crop_size=d_args["crop_size"], 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) if o_args["freeze_weights"]: @@ -226,18 +227,18 @@ def main(args, resume_preempt=False): ) scheduler = None - if o_args.get("use_cosine_schedule", False): + if o_args["use_cosine_schedule"]: 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( crop_size=d_args["crop_size"], - crop_scale=tuple(d_args.get("crop_scale", (0.3, 1.0))), - horizontal_flip=d_args.get("use_horizontal_flip", False), - color_distortion=d_args.get("use_color_distortion", False), - color_jitter=d_args.get("color_jitter_strength", 1.0), - gaussian_blur=d_args.get("use_gaussian_blur", False), + crop_scale=tuple(d_args["crop_scale"]), + horizontal_flip=d_args["use_horizontal_flip"], + color_distortion=d_args["use_color_distortion"], + color_jitter=d_args["color_jitter_strength"], + gaussian_blur=d_args["use_gaussian_blur"], ) val_transform = make_transform_eval( crop_size=d_args["crop_size"], @@ -251,8 +252,8 @@ def main(args, resume_preempt=False): rank=rank, world_size=world_size, collator=None, - num_workers=d_args.get("num_workers", 8), - pin_mem=d_args.get("pin_mem", True), + num_workers=d_args["num_workers"], + pin_mem=d_args["pin_mem"], drop_last=True, ) @@ -264,8 +265,8 @@ def main(args, resume_preempt=False): rank=rank, world_size=world_size, collator=None, - num_workers=d_args.get("num_workers", 8), - pin_mem=d_args.get("pin_mem", True), + num_workers=d_args["num_workers"], + pin_mem=d_args["pin_mem"], drop_last=False, ) @@ -273,11 +274,11 @@ def main(args, resume_preempt=False): model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()]) early_stopper = EarlyStopping( - enabled=es_args.get("enabled", False), - patience=es_args.get("patience", 10), - min_delta=es_args.get("min_delta", 0.0), - min_epochs=es_args.get("min_epochs", 0), - restore_best_weights=es_args.get("restore_best_weights", True), + enabled=es_args["enabled"], + patience=es_args["patience"], + min_delta=es_args["min_delta"], + min_epochs=es_args["min_epochs"], + restore_best_weights=es_args["restore_best_weights"], ) start_epoch = 0 @@ -287,10 +288,16 @@ def main(args, resume_preempt=False): if os.path.exists(latest_path): checkpoint_to_load = latest_path resuming_interrupted = True - elif m_args.get("load_checkpoint", False): - r_file = m_args.get("read_checkpoint") - if r_file is not None: - checkpoint_to_load = os.path.join(folder, r_file) + elif m_args["load_checkpoint"]: + r_file = m_args["read_checkpoint"] + checkpoint_folder = m_args["checkpoint_folder"] + 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): checkpoint = torch.load(checkpoint_to_load, map_location="cpu") @@ -342,7 +349,7 @@ def main(args, resume_preempt=False): if is_best: 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"]): train_sampler.set_epoch(epoch)