From 7e25f2a57fcf92fa598523dbc2ba60b84582ff79 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Fri, 22 May 2026 16:08:47 +0200 Subject: [PATCH] auto_folder --- configs/eval_wilds_vith14.yaml | 2 +- configs/supervised_wilds_vith14_ep300-lp.yaml | 2 +- src/eval_wilds.py | 30 ++++++-- src/train_supervised.py | 34 +++++++++- src/utils/logging.py | 68 +++++++++++++++++++ 5 files changed, 125 insertions(+), 11 deletions(-) diff --git a/configs/eval_wilds_vith14.yaml b/configs/eval_wilds_vith14.yaml index 42d80bf..7c2f8b4 100644 --- a/configs/eval_wilds_vith14.yaml +++ b/configs/eval_wilds_vith14.yaml @@ -18,5 +18,5 @@ data: download: true logging: - folder: experiment_logs/eval-wilds/vith14-ep300-sgd-earlystopping-lr0.01-wd0.0005/ write_tag: iwildcam_test + auto_folder: true diff --git a/configs/supervised_wilds_vith14_ep300-lp.yaml b/configs/supervised_wilds_vith14_ep300-lp.yaml index 5823743..b1ffe61 100644 --- a/configs/supervised_wilds_vith14_ep300-lp.yaml +++ b/configs/supervised_wilds_vith14_ep300-lp.yaml @@ -55,5 +55,5 @@ validation: eval_every: 1 logging: - folder: experiment_logs/supervised/ write_tag: linear_probe + auto_folder: true diff --git a/src/eval_wilds.py b/src/eval_wilds.py index eef04f1..953e782 100644 --- a/src/eval_wilds.py +++ b/src/eval_wilds.py @@ -14,6 +14,7 @@ from src.helper import init_model from src.models.head import ViTClassifier from src.transforms import make_transform_eval from src.utils.distributed import init_distributed +from src.utils.logging import resolve_log_dir _GLOBAL_SEED = 0 @@ -74,9 +75,19 @@ def _load_model_state(model, checkpoint_path, device): def main(args): - world_size, rank = init_distributed() + force_single = bool(args.get("meta", {}).get("force_single_process", False)) + if force_single: + world_size, rank = 1, 0 + else: + world_size, rank = init_distributed() - if dist.is_available() and dist.is_initialized() and world_size > 1 and rank != 0: + if ( + not force_single + and dist.is_available() + and dist.is_initialized() + and world_size > 1 + and rank != 0 + ): dist.barrier() dist.destroy_process_group() return @@ -93,11 +104,13 @@ def main(args): else: device = torch.device(f"cuda:{torch.cuda.current_device()}") - folder = log_args.get("folder", "eval_logs") + folder = resolve_log_dir(args, stage="eval") tag = log_args.get("write_tag", "wilds_eval") - os.makedirs(folder, exist_ok=True) - with open(os.path.join(folder, "params-eval.yaml"), "w") as f: + params_path = os.path.join(folder, "params-eval.yaml") + with open(params_path, "w") as f: + yaml.dump(args, f) + with open(os.path.join(folder, "params.yaml"), "w") as f: yaml.dump(args, f) model_name = meta_args["model_name"] @@ -171,7 +184,12 @@ def main(args): logger.info(f"Eval metrics saved to {metrics_path}") logger.info(f"Eval metrics: {metrics}") - if dist.is_available() and dist.is_initialized() and world_size > 1: + if ( + not force_single + and dist.is_available() + and dist.is_initialized() + and world_size > 1 + ): dist.barrier() dist.destroy_process_group() diff --git a/src/train_supervised.py b/src/train_supervised.py index 8950952..42a4400 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -15,8 +15,9 @@ from src.helper import init_model from src.models.head import ViTClassifier from src.transforms import make_transforms, make_transform_eval from src.utils.distributed import init_distributed -from src.utils.logging import CSVLogger, AverageMeter +from src.utils.logging import CSVLogger, AverageMeter, resolve_log_dir from src.utils.optimizers import LARS +from src.eval_wilds import main as eval_wilds_main # -- log_freq = 10 @@ -175,9 +176,8 @@ def main(args, resume_preempt=False): v_args = args["validation"] es_args = o_args["early_stopping"] - folder = l_args["folder"] + folder = resolve_log_dir(args, stage="train") tag = l_args["write_tag"] - os.makedirs(folder, exist_ok=True) with open(os.path.join(folder, "params-supervised.yaml"), "w") as f: yaml.dump(args, f) @@ -469,6 +469,34 @@ def main(args, resume_preempt=False): best_path, ) + if rank == 0: + eval_args = { + "meta": { + "seed": m_args.get("seed", _GLOBAL_SEED), + "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)), + "crop_size": d_args.get("crop_size", m_args.get("crop_size", 224)), + "use_bfloat16": m_args.get("use_bfloat16", True), + "checkpoint_path": best_path, + "force_single_process": True, + }, + "data": { + "batch_size": d_args.get("batch_size", 128), + "root_path": d_args.get("root_path", "./wilds_data"), + "num_workers": d_args.get("num_workers", 8), + "pin_mem": d_args.get("pin_mem", True), + "split": "test", + "download": True, + }, + "logging": { + "write_tag": "iwildcam_test", + "auto_folder": True, + }, + } + eval_wilds_main(args=eval_args) + if __name__ == "__main__": raise RuntimeError( diff --git a/src/utils/logging.py b/src/utils/logging.py index a252ccb..4a2e794 100644 --- a/src/utils/logging.py +++ b/src/utils/logging.py @@ -5,6 +5,7 @@ # LICENSE file in the root directory of this source tree. # +import os import torch @@ -28,6 +29,73 @@ def gpu_timer(closure, log_timings=True): return result, elapsed_time +def _format_value(value): + if value is None: + return None + if isinstance(value, float): + return f"{value:g}" + return str(value) + + +def build_run_name(args): + meta_args = args.get("meta", {}) + data_args = args.get("data", {}) + opt_args = args.get("optimization", {}) + mask_args = args.get("mask", {}) + + model_name = meta_args.get("model_name", "model") + 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") + optimizer = opt_args.get("optimizer", "opt") + lr = opt_args.get("lr") + weight_decay = opt_args.get("weight_decay") + epochs = opt_args.get("epochs") + + parts = [ + model_name, + f"p{_format_value(patch_size)}" if patch_size is not None else None, + f"c{_format_value(crop_size)}" if crop_size is not None else None, + f"bs{_format_value(batch_size)}" if batch_size is not None else None, + str(optimizer).lower(), + f"lr{_format_value(lr)}" if lr is not None else None, + f"wd{_format_value(weight_decay)}" if weight_decay is not None else None, + f"ep{_format_value(epochs)}" if epochs is not None else None, + ] + return "-".join([p for p in parts if p]) + + +def _extract_run_name_from_checkpoint(meta_args): + checkpoint_path = meta_args.get("checkpoint_path") + if checkpoint_path: + folder = os.path.dirname(checkpoint_path) + else: + folder = meta_args.get("checkpoint_folder") + if not folder: + return None + return os.path.basename(os.path.normpath(folder)) + + +def resolve_log_dir(args, stage="train"): + log_args = args.setdefault("logging", {}) + auto_folder = log_args.get("auto_folder", True) + if auto_folder or not log_args.get("folder"): + run_name = log_args.get("run_name") + if stage == "eval" and not run_name: + meta_args = args.get("meta", {}) + run_name = _extract_run_name_from_checkpoint(meta_args) + if not run_name: + run_name = build_run_name(args) + base_dir = "experiment_logs" + if stage == "eval": + folder = os.path.join(base_dir, "eval-wilds", run_name) + else: + folder = os.path.join(base_dir, run_name) + log_args["folder"] = folder + os.makedirs(log_args["folder"], exist_ok=True) + return log_args["folder"] + + class CSVLogger(object): def __init__(self, fname, *argv):