remove mlp + add batch norm + last 4 avg pool layer as in the paper of IJEPA
This commit is contained in:
@@ -6,9 +6,8 @@ meta:
|
||||
patch_size: 14
|
||||
crop_size: 224
|
||||
use_bfloat16: true
|
||||
probe_type: linear
|
||||
mlp_hidden_dim:
|
||||
dropout: 0.0
|
||||
representation_type: last_avgpool
|
||||
head_type: linear
|
||||
checkpoint_folder: experiment_logs/supervised/
|
||||
read_checkpoint: linear_probe-best.pth.tar
|
||||
|
||||
|
||||
@@ -10,9 +10,8 @@ grid:
|
||||
optimization.momentum: [0.9, 0.99, 0.5]
|
||||
optimization.lr_schedule: [cosine, step]
|
||||
data.batch_size: [32, 64, 128, 256]
|
||||
meta.probe_type: [linear]
|
||||
# meta.mlp_hidden_dim: [512, 1024, 2048]
|
||||
# meta.dropout: [0.0, 0.1, 0.3]
|
||||
meta.representation_type: [last_avgpool]
|
||||
meta.head_type: [linear]
|
||||
|
||||
launch:
|
||||
folder: submitit_logs/
|
||||
|
||||
@@ -6,9 +6,8 @@ meta:
|
||||
read_checkpoint: jepa-ep150.pth.tar
|
||||
use_bfloat16: true
|
||||
num_classes: 182
|
||||
probe_type: linear
|
||||
mlp_hidden_dim: 1024
|
||||
dropout: 0.3
|
||||
representation_type: last_avgpool
|
||||
head_type: linear
|
||||
|
||||
data:
|
||||
batch_size: 128
|
||||
|
||||
@@ -6,9 +6,8 @@ meta:
|
||||
read_checkpoint: jepa-ep300.pth.tar
|
||||
use_bfloat16: true
|
||||
num_classes: 182
|
||||
probe_type: mlp
|
||||
mlp_hidden_dim: 1024
|
||||
dropout: 0.3
|
||||
representation_type: last_avgpool
|
||||
head_type: linear
|
||||
|
||||
data:
|
||||
batch_size: 256
|
||||
|
||||
+8
-3
@@ -131,9 +131,8 @@ def main(args):
|
||||
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),
|
||||
representation_type=meta_args.get("representation_type", "last_avgpool"),
|
||||
head_type=meta_args.get("head_type", "linear"),
|
||||
).to(device)
|
||||
|
||||
checkpoint_path = _resolve_checkpoint_path(meta_args)
|
||||
@@ -200,6 +199,12 @@ def main(args):
|
||||
dist.barrier()
|
||||
dist.destroy_process_group()
|
||||
|
||||
return {
|
||||
"metrics": metrics,
|
||||
"metrics_path": metrics_path,
|
||||
"folder": folder,
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise RuntimeError(
|
||||
|
||||
+32
-22
@@ -8,43 +8,53 @@ class ViTClassifier(nn.Module):
|
||||
encoder,
|
||||
num_classes,
|
||||
embed_dim,
|
||||
probe_type="linear",
|
||||
mlp_hidden_dim=None,
|
||||
dropout=0.0,
|
||||
representation_type="last_avgpool",
|
||||
head_type="linear",
|
||||
):
|
||||
super().__init__()
|
||||
self.encoder = encoder
|
||||
self.representation_type = str(representation_type).lower()
|
||||
|
||||
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'")
|
||||
if self.representation_type == "last_avgpool":
|
||||
in_dim = embed_dim
|
||||
elif self.representation_type == "last4_avgpool_concat":
|
||||
in_dim = 4 * embed_dim
|
||||
else:
|
||||
raise ValueError(f"Unknown representation_type: {representation_type}")
|
||||
|
||||
head_type = str(head_type).lower()
|
||||
if head_type == "linear":
|
||||
self.head = nn.Linear(in_dim, num_classes)
|
||||
elif head_type == "bn_linear":
|
||||
self.head = nn.Sequential(
|
||||
nn.Linear(embed_dim, mlp_hidden_dim),
|
||||
nn.GELU(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(mlp_hidden_dim, num_classes),
|
||||
nn.BatchNorm1d(in_dim, affine=False, eps=1e-6),
|
||||
nn.Linear(in_dim, num_classes),
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unknown probe_type: {probe_type}")
|
||||
raise ValueError(f"Unknown head_type: {head_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)
|
||||
if any(p.requires_grad for p in self.encoder.parameters()):
|
||||
def _extract_representation(self, x):
|
||||
if self.representation_type == "last_avgpool":
|
||||
features = self.encoder(x)
|
||||
return features.mean(dim=1)
|
||||
|
||||
_, layer_outputs = self.encoder(
|
||||
x, return_layer_outputs=True, num_last_layers=4
|
||||
)
|
||||
pooled = [layer_tokens.mean(dim=1) for layer_tokens in layer_outputs]
|
||||
return torch.cat(pooled, dim=-1)
|
||||
|
||||
def forward(self, x):
|
||||
if any(p.requires_grad for p in self.encoder.parameters()):
|
||||
representation = self._extract_representation(x)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
features = self.encoder(x)
|
||||
representation = self._extract_representation(x)
|
||||
|
||||
# Average Pool -> (B, D)
|
||||
avg_embed = features.mean(dim=1)
|
||||
|
||||
logits = self.head(avg_embed)
|
||||
logits = self.head(representation)
|
||||
return logits
|
||||
|
||||
@@ -398,7 +398,13 @@ class VisionTransformer(nn.Module):
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, x, masks=None):
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
masks=None,
|
||||
return_layer_outputs=False,
|
||||
num_last_layers=4,
|
||||
):
|
||||
if masks is not None:
|
||||
if not isinstance(masks, list):
|
||||
masks = [masks]
|
||||
@@ -416,11 +422,20 @@ class VisionTransformer(nn.Module):
|
||||
x = apply_masks(x, masks)
|
||||
|
||||
# -- fwd prop
|
||||
layer_outputs = []
|
||||
for i, blk in enumerate(self.blocks):
|
||||
x = blk(x)
|
||||
if return_layer_outputs:
|
||||
layer_outputs.append(x)
|
||||
|
||||
if self.norm is not None:
|
||||
x = self.norm(x)
|
||||
if return_layer_outputs:
|
||||
layer_outputs = [self.norm(t) for t in layer_outputs]
|
||||
|
||||
if return_layer_outputs:
|
||||
k = min(int(num_last_layers), len(layer_outputs))
|
||||
return x, layer_outputs[-k:]
|
||||
|
||||
return x
|
||||
|
||||
|
||||
+165
-107
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
import json
|
||||
import yaml
|
||||
import logging
|
||||
|
||||
@@ -177,93 +178,14 @@ def main(args, resume_preempt=False):
|
||||
v_args = args["validation"]
|
||||
es_args = o_args["early_stopping"]
|
||||
|
||||
folder = resolve_log_dir(args, stage="train")
|
||||
root_folder = resolve_log_dir(args, stage="train")
|
||||
tag = l_args["write_tag"]
|
||||
|
||||
with open(os.path.join(folder, "params-supervised.yaml"), "w") as f:
|
||||
with open(os.path.join(root_folder, "params-supervised.yaml"), "w") as f:
|
||||
yaml.dump(args, f)
|
||||
with open(os.path.join(folder, "params.yaml"), "w") as f:
|
||||
with open(os.path.join(root_folder, "params.yaml"), "w") as f:
|
||||
yaml.dump(args, f)
|
||||
|
||||
save_path = os.path.join(folder, f"{tag}" + "-ep{epoch}.pth.tar")
|
||||
latest_path = os.path.join(folder, f"{tag}-latest.pth.tar")
|
||||
best_path = os.path.join(folder, f"{tag}-best.pth.tar")
|
||||
log_file = os.path.join(folder, f"{tag}_r{rank}.csv")
|
||||
|
||||
csv_logger = CSVLogger(
|
||||
log_file,
|
||||
("%d", "epoch"),
|
||||
("%.6f", "train_loss"),
|
||||
("%.6f", "val_loss"),
|
||||
("%.6f", "val_acc"),
|
||||
("%.6e", "lr"),
|
||||
("%.6f", "best_val_loss"),
|
||||
("%d", "best_epoch"),
|
||||
("%d", "early_stop"),
|
||||
)
|
||||
|
||||
encoder, _ = init_model(
|
||||
device=device,
|
||||
patch_size=mk_args["patch_size"],
|
||||
crop_size=d_args["crop_size"],
|
||||
model_name=m_args["model_name"],
|
||||
)
|
||||
|
||||
embed_dim = m_args["embed_dim"]
|
||||
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)")
|
||||
for param in model.encoder.parameters():
|
||||
param.requires_grad = False
|
||||
model.encoder.eval()
|
||||
else:
|
||||
logger.info("Training full model (Fine-tuning mode)")
|
||||
|
||||
params = [p for p in model.parameters() if p.requires_grad]
|
||||
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=o_args.get("momentum", 0.9),
|
||||
weight_decay=o_args["weight_decay"],
|
||||
)
|
||||
|
||||
scheduler = None
|
||||
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 lr_schedule == "cosine":
|
||||
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||
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["crop_scale"]),
|
||||
@@ -303,7 +225,110 @@ def main(args, resume_preempt=False):
|
||||
drop_last=False,
|
||||
)
|
||||
|
||||
combinations = [
|
||||
("last_avgpool", "linear"),
|
||||
("last_avgpool", "bn_linear"),
|
||||
("last4_avgpool_concat", "linear"),
|
||||
("last4_avgpool_concat", "bn_linear"),
|
||||
]
|
||||
|
||||
criterion = nn.CrossEntropyLoss().to(device)
|
||||
eval_every = int(v_args["eval_every"])
|
||||
combo_results = []
|
||||
|
||||
for representation_type, head_type in combinations:
|
||||
combo_name = f"{representation_type}__{head_type}"
|
||||
combo_folder = os.path.join(root_folder, combo_name)
|
||||
os.makedirs(combo_folder, exist_ok=True)
|
||||
|
||||
combo_args = yaml.safe_load(yaml.dump(args))
|
||||
combo_args.setdefault("meta", {})["representation_type"] = representation_type
|
||||
combo_args.setdefault("meta", {})["head_type"] = head_type
|
||||
|
||||
with open(os.path.join(combo_folder, "params-supervised.yaml"), "w") as f:
|
||||
yaml.dump(combo_args, f)
|
||||
with open(os.path.join(combo_folder, "params.yaml"), "w") as f:
|
||||
yaml.dump(combo_args, f)
|
||||
|
||||
save_path = os.path.join(combo_folder, f"{tag}" + "-ep{epoch}.pth.tar")
|
||||
latest_path = os.path.join(combo_folder, f"{tag}-latest.pth.tar")
|
||||
best_path = os.path.join(combo_folder, f"{tag}-best.pth.tar")
|
||||
log_file = os.path.join(combo_folder, f"{tag}_r{rank}.csv")
|
||||
|
||||
csv_logger = CSVLogger(
|
||||
log_file,
|
||||
("%d", "epoch"),
|
||||
("%.6f", "train_loss"),
|
||||
("%.6f", "val_loss"),
|
||||
("%.6f", "val_acc"),
|
||||
("%.6e", "lr"),
|
||||
("%.6f", "best_val_loss"),
|
||||
("%d", "best_epoch"),
|
||||
("%d", "early_stop"),
|
||||
)
|
||||
|
||||
encoder, _ = init_model(
|
||||
device=device,
|
||||
patch_size=mk_args["patch_size"],
|
||||
crop_size=d_args["crop_size"],
|
||||
model_name=m_args["model_name"],
|
||||
)
|
||||
|
||||
model = ViTClassifier(
|
||||
encoder,
|
||||
m_args["num_classes"],
|
||||
m_args["embed_dim"],
|
||||
representation_type=representation_type,
|
||||
head_type=head_type,
|
||||
).to(device)
|
||||
|
||||
if o_args["freeze_weights"]:
|
||||
logger.info(
|
||||
f"[{combo_name}] Freezing encoder weights (Linear Probing mode)"
|
||||
)
|
||||
for param in model.encoder.parameters():
|
||||
param.requires_grad = False
|
||||
model.encoder.eval()
|
||||
else:
|
||||
logger.info(f"[{combo_name}] Training full model (Fine-tuning mode)")
|
||||
|
||||
params = [p for p in model.parameters() if p.requires_grad]
|
||||
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=o_args.get("momentum", 0.9),
|
||||
weight_decay=o_args["weight_decay"],
|
||||
)
|
||||
|
||||
scheduler = None
|
||||
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 lr_schedule == "cosine":
|
||||
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||
optimizer, T_max=o_args["epochs"], eta_min=o_args["final_lr"]
|
||||
)
|
||||
|
||||
model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()])
|
||||
|
||||
early_stopper = EarlyStopping(
|
||||
@@ -343,7 +368,7 @@ def main(args, resume_preempt=False):
|
||||
start_epoch = int(checkpoint.get("epoch", 0))
|
||||
early_stopper.load_state_dict(checkpoint.get("early_stopping", {}))
|
||||
logger.info(
|
||||
f"Resuming training from {checkpoint_to_load} at epoch {start_epoch}"
|
||||
f"[{combo_name}] Resuming training from {checkpoint_to_load} at epoch {start_epoch}"
|
||||
)
|
||||
else:
|
||||
encoder_state = checkpoint.get("encoder")
|
||||
@@ -360,7 +385,7 @@ def main(args, resume_preempt=False):
|
||||
encoder_state = strip_module_prefix(encoder_state)
|
||||
msg = model.module.encoder.load_state_dict(encoder_state, strict=False)
|
||||
logger.info(
|
||||
f"Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}"
|
||||
f"[{combo_name}] Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}"
|
||||
)
|
||||
|
||||
def save_checkpoint(epoch, train_loss, val_loss, val_acc, is_best=False):
|
||||
@@ -372,7 +397,7 @@ def main(args, resume_preempt=False):
|
||||
"train_loss": train_loss,
|
||||
"val_loss": val_loss,
|
||||
"val_acc": val_acc,
|
||||
"args": args,
|
||||
"args": combo_args,
|
||||
"early_stopping": early_stopper.state_dict(),
|
||||
}
|
||||
if rank == 0:
|
||||
@@ -382,8 +407,6 @@ def main(args, resume_preempt=False):
|
||||
if is_best:
|
||||
torch.save(save_dict, best_path)
|
||||
|
||||
eval_every = int(v_args["eval_every"])
|
||||
|
||||
for epoch in range(start_epoch, o_args["epochs"]):
|
||||
train_sampler.set_epoch(epoch)
|
||||
val_sampler.set_epoch(epoch)
|
||||
@@ -412,12 +435,14 @@ def main(args, resume_preempt=False):
|
||||
|
||||
if itr % log_freq == 0 and rank == 0:
|
||||
logger.info(
|
||||
f"Epoch {epoch + 1} [{itr}/{len(train_loader)}] Train Loss: {loss_meter.avg:.4f}"
|
||||
f"[{combo_name}] Epoch {epoch + 1} [{itr}/{len(train_loader)}] Train Loss: {loss_meter.avg:.4f}"
|
||||
)
|
||||
|
||||
train_loss = distributed_average(loss_meter.avg, device)
|
||||
|
||||
do_eval = ((epoch + 1) % eval_every == 0) or (epoch + 1 == o_args["epochs"])
|
||||
do_eval = ((epoch + 1) % eval_every == 0) or (
|
||||
epoch + 1 == o_args["epochs"]
|
||||
)
|
||||
if do_eval:
|
||||
val_loss, val_acc = evaluate(
|
||||
model=model,
|
||||
@@ -426,7 +451,9 @@ def main(args, resume_preempt=False):
|
||||
device=device,
|
||||
use_bfloat16=m_args["use_bfloat16"],
|
||||
)
|
||||
is_best, should_stop = early_stopper.step(epoch + 1, val_loss, model.module)
|
||||
is_best, should_stop = early_stopper.step(
|
||||
epoch + 1, val_loss, model.module
|
||||
)
|
||||
else:
|
||||
val_loss = float("nan")
|
||||
val_acc = float("nan")
|
||||
@@ -437,7 +464,7 @@ def main(args, resume_preempt=False):
|
||||
|
||||
if rank == 0:
|
||||
logger.info(
|
||||
f"Epoch {epoch + 1} done | train_loss={train_loss:.6f} val_loss={val_loss:.6f} val_acc={val_acc:.6f} best_val_loss={early_stopper.best_metric:.6f}"
|
||||
f"[{combo_name}] Epoch {epoch + 1} done | train_loss={train_loss:.6f} val_loss={val_loss:.6f} val_acc={val_acc:.6f} best_val_loss={early_stopper.best_metric:.6f}"
|
||||
)
|
||||
|
||||
csv_logger.log(
|
||||
@@ -454,19 +481,23 @@ def main(args, resume_preempt=False):
|
||||
save_checkpoint(epoch + 1, train_loss, val_loss, val_acc, is_best=is_best)
|
||||
|
||||
stop_tensor = torch.tensor([int(should_stop)], device=device)
|
||||
if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1:
|
||||
if (
|
||||
dist.is_available()
|
||||
and dist.is_initialized()
|
||||
and dist.get_world_size() > 1
|
||||
):
|
||||
dist.broadcast(stop_tensor, src=0)
|
||||
|
||||
if bool(stop_tensor.item()):
|
||||
if rank == 0:
|
||||
logger.info(
|
||||
f"Early stopping at epoch {epoch + 1}. Best val_loss={early_stopper.best_metric:.6f} @ epoch {early_stopper.best_epoch}"
|
||||
f"[{combo_name}] Early stopping at epoch {epoch + 1}. Best val_loss={early_stopper.best_metric:.6f} @ epoch {early_stopper.best_epoch}"
|
||||
)
|
||||
break
|
||||
|
||||
if early_stopper.enabled and early_stopper.restore_best_weights:
|
||||
if rank == 0:
|
||||
logger.info("Restoring best model weights before exit")
|
||||
logger.info(f"[{combo_name}] Restoring best model weights before exit")
|
||||
early_stopper.restore(model.module, device)
|
||||
if rank == 0:
|
||||
torch.save(
|
||||
@@ -474,11 +505,21 @@ def main(args, resume_preempt=False):
|
||||
"model": model.module.state_dict(),
|
||||
"epoch": early_stopper.best_epoch,
|
||||
"val_loss": early_stopper.best_metric,
|
||||
"args": args,
|
||||
"args": combo_args,
|
||||
},
|
||||
best_path,
|
||||
)
|
||||
|
||||
combo_result = {
|
||||
"representation_type": representation_type,
|
||||
"head_type": head_type,
|
||||
"combo_name": combo_name,
|
||||
"best_val_loss": float(early_stopper.best_metric),
|
||||
"best_epoch": int(early_stopper.best_epoch),
|
||||
"best_checkpoint": best_path,
|
||||
"train_folder": combo_folder,
|
||||
}
|
||||
|
||||
if rank == 0:
|
||||
eval_args = {
|
||||
"meta": {
|
||||
@@ -486,12 +527,13 @@ def main(args, resume_preempt=False):
|
||||
"model_name": m_args["model_name"],
|
||||
"embed_dim": m_args["embed_dim"],
|
||||
"num_classes": m_args["num_classes"],
|
||||
"patch_size": mk_args.get("patch_size", m_args.get("patch_size", 16)),
|
||||
"patch_size": mk_args.get(
|
||||
"patch_size", m_args.get("patch_size", 16)
|
||||
),
|
||||
"crop_size": d_args.get("crop_size", m_args.get("crop_size", 224)),
|
||||
"use_bfloat16": m_args.get("use_bfloat16", True),
|
||||
"probe_type": m_args.get("probe_type", "linear"),
|
||||
"mlp_hidden_dim": m_args.get("mlp_hidden_dim"),
|
||||
"dropout": m_args.get("dropout", 0.0),
|
||||
"representation_type": representation_type,
|
||||
"head_type": head_type,
|
||||
"checkpoint_path": best_path,
|
||||
"force_single_process": True,
|
||||
},
|
||||
@@ -508,10 +550,9 @@ def main(args, resume_preempt=False):
|
||||
"auto_folder": True,
|
||||
},
|
||||
}
|
||||
eval_wilds_main(args=eval_args)
|
||||
|
||||
eval_result = eval_wilds_main(args=eval_args)
|
||||
eval_folder = eval_args.get("logging", {}).get("folder")
|
||||
supervised_params = os.path.join(folder, "params-supervised.yaml")
|
||||
supervised_params = os.path.join(combo_folder, "params-supervised.yaml")
|
||||
if eval_folder and os.path.exists(supervised_params):
|
||||
try:
|
||||
shutil.copy2(
|
||||
@@ -525,11 +566,28 @@ def main(args, resume_preempt=False):
|
||||
except OSError:
|
||||
logger.warning("Could not copy supervised params to eval folder")
|
||||
|
||||
if os.path.exists(folder):
|
||||
try:
|
||||
shutil.rmtree(folder)
|
||||
except OSError:
|
||||
logger.warning("Could not remove supervised run folder")
|
||||
combo_result["eval"] = eval_result
|
||||
|
||||
combo_results.append(combo_result)
|
||||
|
||||
if (
|
||||
dist.is_available()
|
||||
and dist.is_initialized()
|
||||
and dist.get_world_size() > 1
|
||||
):
|
||||
dist.barrier()
|
||||
|
||||
if rank == 0:
|
||||
best_combo = min(combo_results, key=lambda r: r["best_val_loss"])
|
||||
summary = {
|
||||
"protocol": "ijepa_linear_probe_best_of_4",
|
||||
"combinations": combo_results,
|
||||
"best_by_val_loss": best_combo,
|
||||
}
|
||||
summary_path = os.path.join(root_folder, "ijepa_linear_probe_summary.json")
|
||||
with open(summary_path, "w") as f:
|
||||
json.dump(summary, f, indent=2, sort_keys=True)
|
||||
logger.info(f"Saved IJEPA LP summary to {summary_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user