diff --git a/configs/grids/lp_grid.yaml b/configs/grids/lp_grid.yaml index cae0969..091ced0 100644 --- a/configs/grids/lp_grid.yaml +++ b/configs/grids/lp_grid.yaml @@ -4,12 +4,15 @@ constants: logging.write_tag: linear_probe grid: - optimization.optimizer: [adamw] - optimization.lr: [0.01, 0.05, 0.001] - optimization.weight_decay: [0.0, 5.0e-4] + optimization.optimizer: [adamw, lars, sgd] + optimization.lr: [0.01, 0.05] + optimization.weight_decay: [5.0e-4] optimization.momentum: [0.9, 0.5] - optimization.lr_schedule: [step, cosine] - data.batch_size: [256, 128] + optimization.lr_schedule: [cosine, step] + data.batch_size: [32, 64, 512] + meta.probe_type: [linear, mlp] + meta.mlp_hidden_dim: [1024] + meta.dropout: [0.0] launch: folder: submitit_logs/ diff --git a/configs/supervised_wilds_vith14_ep300-lp.yaml b/configs/supervised_wilds_vith14_ep300-lp.yaml index 1f1ff7d..bac988f 100644 --- a/configs/supervised_wilds_vith14_ep300-lp.yaml +++ b/configs/supervised_wilds_vith14_ep300-lp.yaml @@ -6,6 +6,9 @@ meta: read_checkpoint: jepa-ep300.pth.tar use_bfloat16: true num_classes: 182 + probe_type: linear + mlp_hidden_dim: 1024 + dropout: 0.0 data: batch_size: 256 diff --git a/src/eval_wilds.py b/src/eval_wilds.py index 953e782..7b88d57 100644 --- a/src/eval_wilds.py +++ b/src/eval_wilds.py @@ -127,7 +127,14 @@ def main(args): crop_size=crop_size, model_name=model_name, ) - model = ViTClassifier(encoder, num_classes, embed_dim).to(device) + model = ViTClassifier( + encoder, + num_classes, + embed_dim, + probe_type=meta_args.get("probe_type", "linear"), + mlp_hidden_dim=meta_args.get("mlp_hidden_dim"), + dropout=meta_args.get("dropout", 0.0), + ).to(device) checkpoint_path = _resolve_checkpoint_path(meta_args) _load_model_state(model, checkpoint_path, device) diff --git a/src/models/head.py b/src/models/head.py index 4f17237..3a7785e 100644 --- a/src/models/head.py +++ b/src/models/head.py @@ -3,13 +3,37 @@ import torch.nn as nn class ViTClassifier(nn.Module): - def __init__(self, encoder, num_classes, embed_dim): + def __init__( + self, + encoder, + num_classes, + embed_dim, + probe_type="linear", + mlp_hidden_dim=None, + dropout=0.0, + ): super().__init__() self.encoder = encoder - self.head = nn.Linear(embed_dim, num_classes) - nn.init.trunc_normal_(self.head.weight, std=0.01) - nn.init.zeros_(self.head.bias) + probe_type = str(probe_type).lower() + if probe_type == "linear": + self.head = nn.Linear(embed_dim, num_classes) + elif probe_type == "mlp": + if mlp_hidden_dim is None: + raise ValueError("mlp_hidden_dim must be set for probe_type='mlp'") + self.head = nn.Sequential( + nn.Linear(embed_dim, mlp_hidden_dim), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(mlp_hidden_dim, num_classes), + ) + else: + raise ValueError(f"Unknown probe_type: {probe_type}") + + for module in self.head.modules(): + if isinstance(module, nn.Linear): + nn.init.trunc_normal_(module.weight, std=0.01) + nn.init.zeros_(module.bias) def forward(self, x): # ViT -> (B, N, D) diff --git a/src/train_supervised.py b/src/train_supervised.py index 1b60ed9..a356626 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -207,7 +207,14 @@ def main(args, resume_preempt=False): ) embed_dim = m_args["embed_dim"] - model = ViTClassifier(encoder, m_args["num_classes"], embed_dim).to(device) + model = ViTClassifier( + encoder, + m_args["num_classes"], + embed_dim, + probe_type=m_args.get("probe_type", "linear"), + mlp_hidden_dim=m_args.get("mlp_hidden_dim"), + dropout=m_args.get("dropout", 0.0), + ).to(device) if o_args["freeze_weights"]: logger.info("Freezing encoder weights (Linear Probing mode)")