fix: do not run 4 times and log new params

This commit is contained in:
YannAhlgrim
2026-06-02 12:00:12 +02:00
parent 2556c17d88
commit 150c2a7253
3 changed files with 360 additions and 347 deletions
+2 -2
View File
@@ -10,8 +10,8 @@ grid:
optimization.momentum: [0.9, 0.99, 0.5]
optimization.lr_schedule: [cosine, step]
data.batch_size: [32, 64, 128, 256]
meta.representation_type: [last_avgpool]
meta.head_type: [linear]
meta.representation_type: [last_avgpool, last4_avgpool_concat]
meta.head_type: [linear, bn_linear]
launch:
folder: submitit_logs/
+181 -172
View File
@@ -59,6 +59,42 @@ def distributed_sum(value, device):
return float(tensor.item())
def _find_metric(metrics, key):
if isinstance(metrics, dict):
if key in metrics:
return metrics[key]
for value in metrics.values():
found = _find_metric(value, key)
if found is not None:
return found
if isinstance(metrics, list):
for item in metrics:
found = _find_metric(item, key)
if found is not None:
return found
return None
def _collect_eval_rows(eval_root, metric_key):
rows = []
for dirpath, _, filenames in os.walk(eval_root):
for fname in filenames:
if not fname.endswith("_metrics.json"):
continue
path = os.path.join(dirpath, fname)
try:
with open(path, "r") as f:
metrics = json.load(f)
except (OSError, json.JSONDecodeError):
continue
value = _find_metric(metrics, metric_key)
if value is None:
continue
run_name = os.path.basename(os.path.dirname(path))
rows.append((float(value), run_name, path))
return rows
class EarlyStopping:
def __init__(
self,
@@ -178,14 +214,93 @@ def main(args, resume_preempt=False):
v_args = args["validation"]
es_args = o_args["early_stopping"]
root_folder = resolve_log_dir(args, stage="train")
folder = resolve_log_dir(args, stage="train")
tag = l_args["write_tag"]
with open(os.path.join(root_folder, "params-supervised.yaml"), "w") as f:
with open(os.path.join(folder, "params-supervised.yaml"), "w") as f:
yaml.dump(args, f)
with open(os.path.join(root_folder, "params.yaml"), "w") as f:
with open(os.path.join(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"],
)
representation_type = m_args.get("representation_type", "last_avgpool")
head_type = m_args.get("head_type", "linear")
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("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"]),
@@ -225,112 +340,11 @@ 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()])
eval_every = int(v_args["eval_every"])
early_stopper = EarlyStopping(
enabled=es_args["enabled"],
patience=es_args["patience"],
@@ -368,7 +382,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"[{combo_name}] Resuming training from {checkpoint_to_load} at epoch {start_epoch}"
f"Resuming training from {checkpoint_to_load} at epoch {start_epoch}"
)
else:
encoder_state = checkpoint.get("encoder")
@@ -385,7 +399,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"[{combo_name}] Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}"
f"Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}"
)
def save_checkpoint(epoch, train_loss, val_loss, val_acc, is_best=False):
@@ -397,7 +411,7 @@ def main(args, resume_preempt=False):
"train_loss": train_loss,
"val_loss": val_loss,
"val_acc": val_acc,
"args": combo_args,
"args": args,
"early_stopping": early_stopper.state_dict(),
}
if rank == 0:
@@ -435,14 +449,12 @@ def main(args, resume_preempt=False):
if itr % log_freq == 0 and rank == 0:
logger.info(
f"[{combo_name}] Epoch {epoch + 1} [{itr}/{len(train_loader)}] Train Loss: {loss_meter.avg:.4f}"
f"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,
@@ -451,9 +463,7 @@ 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")
@@ -464,7 +474,7 @@ def main(args, resume_preempt=False):
if rank == 0:
logger.info(
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}"
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}"
)
csv_logger.log(
@@ -481,23 +491,19 @@ 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"[{combo_name}] Early stopping at epoch {epoch + 1}. Best val_loss={early_stopper.best_metric:.6f} @ epoch {early_stopper.best_epoch}"
f"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(f"[{combo_name}] Restoring best model weights before exit")
logger.info("Restoring best model weights before exit")
early_stopper.restore(model.module, device)
if rank == 0:
torch.save(
@@ -505,21 +511,11 @@ def main(args, resume_preempt=False):
"model": model.module.state_dict(),
"epoch": early_stopper.best_epoch,
"val_loss": early_stopper.best_metric,
"args": combo_args,
"args": 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": {
@@ -527,9 +523,7 @@ 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),
"representation_type": representation_type,
@@ -551,43 +545,58 @@ def main(args, resume_preempt=False):
},
}
eval_result = eval_wilds_main(args=eval_args)
eval_folder = eval_args.get("logging", {}).get("folder")
supervised_params = os.path.join(combo_folder, "params-supervised.yaml")
if eval_folder and os.path.exists(supervised_params):
if eval_folder:
try:
shutil.copy2(
supervised_params,
os.path.join(eval_folder, "params-supervised.yaml"),
)
shutil.copy2(
supervised_params,
os.path.join(eval_folder, "params.yaml"),
)
except OSError:
logger.warning("Could not copy supervised params to eval 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,
params_out = yaml.safe_load(yaml.dump(args))
params_out.setdefault("meta", {})["representation_type"] = representation_type
params_out.setdefault("meta", {})["head_type"] = head_type
metric_key = m_args.get("selection_metric", "macro_f1")
eval_root = os.path.join("experiment_logs", "eval-wilds")
ranking = {
"metric_key": metric_key,
"this_run_name": os.path.basename(os.path.normpath(eval_folder)),
"this_is_best": None,
"this_metric": None,
"best_run_name": None,
"best_metric": None,
}
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 os.path.isdir(eval_root):
rows = _collect_eval_rows(eval_root, metric_key)
if rows:
rows.sort(key=lambda r: r[0], reverse=True)
best_value, best_run_name, _ = rows[0]
this_run_name = ranking["this_run_name"]
this_value = None
for value, run_name, _ in rows:
if run_name == this_run_name:
this_value = value
break
ranking["this_metric"] = this_value
ranking["best_run_name"] = best_run_name
ranking["best_metric"] = best_value
if this_value is not None:
ranking["this_is_best"] = this_run_name == best_run_name
params_out["results"] = {
"best_val_loss": float(early_stopper.best_metric),
"best_epoch": int(early_stopper.best_epoch),
"best_checkpoint": best_path,
"eval_metrics": eval_result.get("metrics") if eval_result else None,
"ranking": ranking,
}
with open(os.path.join(eval_folder, "params-supervised.yaml"), "w") as f:
yaml.dump(params_out, f)
with open(os.path.join(eval_folder, "params.yaml"), "w") as f:
yaml.dump(params_out, f)
except OSError:
logger.warning("Could not write supervised params to eval folder")
if os.path.exists(folder):
try:
shutil.rmtree(folder)
except OSError:
logger.warning("Could not remove supervised run folder")
if __name__ == "__main__":
+4
View File
@@ -49,6 +49,8 @@ def build_run_name(args):
val_args = args.get("validation", {})
model_name = meta_args.get("model_name", "model")
representation_type = meta_args.get("representation_type")
head_type = meta_args.get("head_type")
patch_size = mask_args.get("patch_size", meta_args.get("patch_size"))
crop_size = data_args.get("crop_size", meta_args.get("crop_size"))
batch_size = data_args.get("batch_size")
@@ -67,6 +69,8 @@ def build_run_name(args):
add("p", patch_size)
add("c", crop_size)
add("bs", batch_size)
add("rep", representation_type)
add("head", head_type)
parts.append(str(optimizer).lower())
add("lr", lr)
add("wd", weight_decay)