disable random resize crop + ft config
This commit is contained in:
@@ -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
|
||||||
+2
-1
@@ -14,6 +14,7 @@ data:
|
|||||||
pin_mem: true
|
pin_mem: true
|
||||||
crop_size: 224
|
crop_size: 224
|
||||||
crop_scale: [0.3, 1.0]
|
crop_scale: [0.3, 1.0]
|
||||||
|
use_random_resized_crop: false
|
||||||
use_horizontal_flip: false
|
use_horizontal_flip: false
|
||||||
use_color_distortion: false
|
use_color_distortion: false
|
||||||
color_jitter_strength: 0.0
|
color_jitter_strength: 0.0
|
||||||
@@ -44,5 +45,5 @@ validation:
|
|||||||
eval_every: 1
|
eval_every: 1
|
||||||
|
|
||||||
logging:
|
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
|
write_tag: linear_probe
|
||||||
@@ -78,6 +78,7 @@ def main(args, resume_preempt=False):
|
|||||||
use_horizontal_flip = args["data"]["use_horizontal_flip"]
|
use_horizontal_flip = args["data"]["use_horizontal_flip"]
|
||||||
use_color_distortion = args["data"]["use_color_distortion"]
|
use_color_distortion = args["data"]["use_color_distortion"]
|
||||||
color_jitter = args["data"]["color_jitter_strength"]
|
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"]
|
batch_size = args["data"]["batch_size"]
|
||||||
pin_mem = args["data"]["pin_mem"]
|
pin_mem = args["data"]["pin_mem"]
|
||||||
@@ -181,6 +182,7 @@ def main(args, resume_preempt=False):
|
|||||||
horizontal_flip=use_horizontal_flip,
|
horizontal_flip=use_horizontal_flip,
|
||||||
color_distortion=use_color_distortion,
|
color_distortion=use_color_distortion,
|
||||||
color_jitter=color_jitter,
|
color_jitter=color_jitter,
|
||||||
|
use_random_resized_crop=use_random_resized_crop,
|
||||||
)
|
)
|
||||||
|
|
||||||
# -- init data-loaders/samplers
|
# -- init data-loaders/samplers
|
||||||
|
|||||||
@@ -239,6 +239,7 @@ def main(args, resume_preempt=False):
|
|||||||
color_distortion=d_args["use_color_distortion"],
|
color_distortion=d_args["use_color_distortion"],
|
||||||
color_jitter=d_args["color_jitter_strength"],
|
color_jitter=d_args["color_jitter_strength"],
|
||||||
gaussian_blur=d_args["use_gaussian_blur"],
|
gaussian_blur=d_args["use_gaussian_blur"],
|
||||||
|
use_random_resized_crop=d_args.get("use_random_resized_crop", True),
|
||||||
)
|
)
|
||||||
val_transform = make_transform_eval(
|
val_transform = make_transform_eval(
|
||||||
crop_size=d_args["crop_size"],
|
crop_size=d_args["crop_size"],
|
||||||
|
|||||||
+9
-1
@@ -23,6 +23,7 @@ def make_transforms(
|
|||||||
horizontal_flip=False,
|
horizontal_flip=False,
|
||||||
color_distortion=False,
|
color_distortion=False,
|
||||||
gaussian_blur=False,
|
gaussian_blur=False,
|
||||||
|
use_random_resized_crop=True,
|
||||||
normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
|
normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
|
||||||
):
|
):
|
||||||
logger.info("making imagenet data transforms")
|
logger.info("making imagenet data transforms")
|
||||||
@@ -36,7 +37,14 @@ def make_transforms(
|
|||||||
return color_distort
|
return color_distort
|
||||||
|
|
||||||
transform_list = []
|
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:
|
if horizontal_flip:
|
||||||
transform_list += [transforms.RandomHorizontalFlip()]
|
transform_list += [transforms.RandomHorizontalFlip()]
|
||||||
if color_distortion:
|
if color_distortion:
|
||||||
|
|||||||
Reference in New Issue
Block a user