diff --git a/configs/supervised_wilds_vith14_ep300-ft.yaml b/configs/supervised_wilds_vith14_ep300-ft.yaml new file mode 100644 index 0000000..e4d145f --- /dev/null +++ b/configs/supervised_wilds_vith14_ep300-ft.yaml @@ -0,0 +1,49 @@ +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: 128 + root_path: ./wilds_data + num_workers: 10 + pin_mem: true + crop_size: 224 + crop_scale: [0.3, 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: adamw # 'adamw' or 'sgd' + freeze_weights: false # 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 + ipe_scale: 1.0 + early_stopping: + enabled: true + patience: 10 + min_delta: 1.0e-4 + min_epochs: 15 + restore_best_weights: true + +validation: + eval_every: 1 + +logging: + folder: experiment_logs/supervised-vith14.224-bs.128-ep.300-ft/ + write_tag: fine_tune diff --git a/configs/supervised_wilds_vith14_ep300.yaml b/configs/supervised_wilds_vith14_ep300-lp.yaml similarity index 90% rename from configs/supervised_wilds_vith14_ep300.yaml rename to configs/supervised_wilds_vith14_ep300-lp.yaml index e01840e..e1ea6cc 100644 --- a/configs/supervised_wilds_vith14_ep300.yaml +++ b/configs/supervised_wilds_vith14_ep300-lp.yaml @@ -14,6 +14,7 @@ data: pin_mem: true crop_size: 224 crop_scale: [0.3, 1.0] + use_random_resized_crop: false use_horizontal_flip: false use_color_distortion: false color_jitter_strength: 0.0 @@ -44,5 +45,5 @@ validation: eval_every: 1 logging: - folder: experiment_logs/supervised-vith14.224-bs.128-ep.300/ + folder: experiment_logs/supervised-vith14.224-bs.128-ep.300-lp/ write_tag: linear_probe diff --git a/src/train.py b/src/train.py index a1e538a..19dde24 100644 --- a/src/train.py +++ b/src/train.py @@ -78,6 +78,7 @@ def main(args, resume_preempt=False): use_horizontal_flip = args["data"]["use_horizontal_flip"] use_color_distortion = args["data"]["use_color_distortion"] color_jitter = args["data"]["color_jitter_strength"] + use_random_resized_crop = args["data"].get("use_random_resized_crop", True) # -- batch_size = args["data"]["batch_size"] pin_mem = args["data"]["pin_mem"] @@ -181,6 +182,7 @@ def main(args, resume_preempt=False): horizontal_flip=use_horizontal_flip, color_distortion=use_color_distortion, color_jitter=color_jitter, + use_random_resized_crop=use_random_resized_crop, ) # -- init data-loaders/samplers diff --git a/src/train_supervised.py b/src/train_supervised.py index 94c3b58..8b49286 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -239,6 +239,7 @@ def main(args, resume_preempt=False): color_distortion=d_args["use_color_distortion"], color_jitter=d_args["color_jitter_strength"], gaussian_blur=d_args["use_gaussian_blur"], + use_random_resized_crop=d_args.get("use_random_resized_crop", True), ) val_transform = make_transform_eval( crop_size=d_args["crop_size"], diff --git a/src/transforms.py b/src/transforms.py index 739f45d..25f5c45 100644 --- a/src/transforms.py +++ b/src/transforms.py @@ -23,6 +23,7 @@ def make_transforms( horizontal_flip=False, color_distortion=False, gaussian_blur=False, + use_random_resized_crop=True, normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ): logger.info("making imagenet data transforms") @@ -36,7 +37,14 @@ def make_transforms( return color_distort transform_list = [] - transform_list += [transforms.RandomResizedCrop(crop_size, scale=crop_scale)] + if use_random_resized_crop: + transform_list += [transforms.RandomResizedCrop(crop_size, scale=crop_scale)] + else: + resize_size = int(crop_size * 256 / 224) + transform_list += [ + transforms.Resize(resize_size), + transforms.CenterCrop(crop_size), + ] if horizontal_flip: transform_list += [transforms.RandomHorizontalFlip()] if color_distortion: