From d22a3aba90a79ed7ee1f1eccfd40c16dd5bf7151 Mon Sep 17 00:00:00 2001 From: YannAhlgrim Date: Thu, 14 May 2026 10:52:40 +0200 Subject: [PATCH] evaluation pipeline --- configs/eval_wilds_vith14.yaml | 22 ++++ main_eval_wilds.py | 102 ++++++++++++++++++ src/datasets/wilds.py | 29 ++++++ src/eval_wilds.py | 182 +++++++++++++++++++++++++++++++++ 4 files changed, 335 insertions(+) create mode 100644 configs/eval_wilds_vith14.yaml create mode 100644 main_eval_wilds.py create mode 100644 src/eval_wilds.py diff --git a/configs/eval_wilds_vith14.yaml b/configs/eval_wilds_vith14.yaml new file mode 100644 index 0000000..5e529b2 --- /dev/null +++ b/configs/eval_wilds_vith14.yaml @@ -0,0 +1,22 @@ +meta: + seed: 0 + model_name: vit_huge + embed_dim: 1280 + num_classes: 182 + patch_size: 14 + crop_size: 224 + use_bfloat16: true + checkpoint_folder: experiment_logs/supervised-vith14.224-bs.128-ep.300/ + read_checkpoint: linear_probe-best.pth.tar + +data: + batch_size: 128 + root_path: ./wilds_data + num_workers: 10 + pin_mem: true + split: test + download: true + +logging: + folder: experiment_logs/eval-wilds-vith14/ + write_tag: iwildcam_test diff --git a/main_eval_wilds.py b/main_eval_wilds.py new file mode 100644 index 0000000..ed5695d --- /dev/null +++ b/main_eval_wilds.py @@ -0,0 +1,102 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +# + +import argparse +import logging +import os +import pprint +import sys +import yaml + +import submitit + +from src.eval_wilds import main as app_main + +logging.basicConfig(stream=sys.stdout, level=logging.INFO) +logger = logging.getLogger() + + +parser = argparse.ArgumentParser() +parser.add_argument("--folder", type=str, help="location to save submitit logs") +parser.add_argument( + "--batch-launch", + action="store_true", + help="whether fname points to a file to batch-lauch several config files", +) +parser.add_argument( + "--fname", + type=str, + help="yaml file containing config file names to launch", + default="configs.yaml", +) +parser.add_argument("--partition", type=str, help="cluster partition to submit jobs on") +parser.add_argument( + "--nodes", type=int, default=1, help="num. nodes to request for job" +) +parser.add_argument( + "--tasks-per-node", type=int, default=1, help="num. procs to per node" +) +parser.add_argument("--time", type=int, default=4300, help="time in minutes to run job") + + +class Trainer: + def __init__(self, fname="configs.yaml"): + self.fname = fname + + def __call__(self): + fname = self.fname + logger.info(f"called-params {fname}") + + params = None + with open(fname, "r") as y_file: + params = yaml.load(y_file, Loader=yaml.FullLoader) + logger.info("loaded params...") + pp = pprint.PrettyPrinter(indent=4) + pp.pprint(params) + + app_main(args=params) + + def checkpoint(self): + fb_trainer = Trainer(self.fname) + return submitit.helpers.DelayedSubmission( + fb_trainer, + ) + + +def launch(): + executor = submitit.SlurmExecutor( + folder=os.path.join(args.folder, "job_%j"), max_num_timeout=20 + ) + executor.update_parameters( + partition=args.partition, + mem_per_gpu="55G", + time=args.time, + nodes=args.nodes, + ntasks_per_node=args.tasks_per_node, + cpus_per_task=10, + gpus_per_node=args.tasks_per_node, + ) + + config_fnames = [args.fname] + + jobs, trainers = [], [] + with executor.batch(): + for cf in config_fnames: + fb_trainer = Trainer(cf) + job = executor.submit( + fb_trainer, + ) + trainers.append(fb_trainer) + jobs.append(job) + + for job in jobs: + print(job.job_id) + + +if __name__ == "__main__": + args = parser.parse_args() + launch() diff --git a/src/datasets/wilds.py b/src/datasets/wilds.py index c9ab5aa..a6f34d2 100644 --- a/src/datasets/wilds.py +++ b/src/datasets/wilds.py @@ -1,6 +1,7 @@ import torch from logging import getLogger from wilds import get_dataset +from wilds.common.data_loaders import get_eval_loader logger = getLogger() @@ -50,6 +51,34 @@ def make_iwildcam( return dataset, data_loader, dist_sampler +def make_iwildcam_eval( + transform, + batch_size, + split="test", + num_workers=8, + root_path="./wilds_data", + download=True, + pin_mem=True, +): + full_dataset = get_dataset( + dataset="iwildcam", download=download, root_dir=root_path + ) + + dataset = full_dataset.get_subset(split, transform=transform) + + logger.info(f"iWildCam {split} eval dataset created with {len(dataset)} samples") + + data_loader = get_eval_loader( + "standard", + dataset, + batch_size=batch_size, + num_workers=num_workers, + pin_memory=pin_mem, + ) + + return full_dataset, dataset, data_loader + + class WildsToTorchWrapper(torch.utils.data.Dataset): """ Mimics the ImageNet wrapper by always returning (image, target). diff --git a/src/eval_wilds.py b/src/eval_wilds.py new file mode 100644 index 0000000..eef04f1 --- /dev/null +++ b/src/eval_wilds.py @@ -0,0 +1,182 @@ +import os +import sys +import json +import yaml +import logging + +import numpy as np + +import torch +import torch.distributed as dist + +from src.datasets.wilds import make_iwildcam_eval +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 + + +_GLOBAL_SEED = 0 +np.random.seed(_GLOBAL_SEED) +torch.manual_seed(_GLOBAL_SEED) +torch.backends.cudnn.benchmark = True + +logging.basicConfig(stream=sys.stdout, level=logging.INFO) +logger = logging.getLogger() + + +def strip_module_prefix(state_dict): + if not any(k.startswith("module.") for k in state_dict.keys()): + return state_dict + return { + k[len("module.") :] if k.startswith("module.") else k: v + for k, v in state_dict.items() + } + + +def _load_yaml(path): + with open(path, "r") as f: + return yaml.load(f, Loader=yaml.FullLoader) + + +def _get_seed(args): + return int(args.get("meta", {}).get("seed", _GLOBAL_SEED)) + + +def _set_seed(seed): + np.random.seed(seed) + torch.manual_seed(seed) + + +def _resolve_checkpoint_path(meta_args): + if meta_args.get("checkpoint_path"): + return meta_args["checkpoint_path"] + folder = meta_args.get("checkpoint_folder") + fname = meta_args.get("read_checkpoint") + if folder is None or fname is None: + return None + return fname if os.path.isabs(fname) else os.path.join(folder, fname) + + +def _load_model_state(model, checkpoint_path, device): + if checkpoint_path is None: + raise ValueError("checkpoint_path or checkpoint_folder/read_checkpoint required") + if not os.path.exists(checkpoint_path): + raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") + + checkpoint = torch.load(checkpoint_path, map_location="cpu") + if "model" not in checkpoint: + raise KeyError(f"No model weights found in checkpoint: {checkpoint_path}") + state = strip_module_prefix(checkpoint["model"]) + msg = model.load_state_dict(state, strict=True) + model.to(device) + logger.info(f"Loaded model from {checkpoint_path} with msg: {msg}") + + +def main(args): + world_size, rank = init_distributed() + + if dist.is_available() and dist.is_initialized() and world_size > 1 and rank != 0: + dist.barrier() + dist.destroy_process_group() + return + + seed = _get_seed(args) + _set_seed(seed) + + meta_args = args.get("meta", {}) + data_args = args.get("data", {}) + log_args = args.get("logging", {}) + + if not torch.cuda.is_available(): + device = torch.device("cpu") + else: + device = torch.device(f"cuda:{torch.cuda.current_device()}") + + folder = log_args.get("folder", "eval_logs") + 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: + yaml.dump(args, f) + + model_name = meta_args["model_name"] + embed_dim = meta_args["embed_dim"] + num_classes = meta_args["num_classes"] + patch_size = meta_args.get("patch_size", 16) + crop_size = meta_args.get("crop_size", 224) + use_bfloat16 = bool(meta_args.get("use_bfloat16", True)) + use_autocast = use_bfloat16 and torch.cuda.is_available() + + encoder, _ = init_model( + device=device, + patch_size=patch_size, + crop_size=crop_size, + model_name=model_name, + ) + model = ViTClassifier(encoder, num_classes, embed_dim).to(device) + + checkpoint_path = _resolve_checkpoint_path(meta_args) + _load_model_state(model, checkpoint_path, device) + model.eval() + + eval_transform = make_transform_eval(crop_size=crop_size) + split = data_args.get("split", "test") + root_path = data_args.get("root_path", "./wilds_data") + + full_dataset, eval_dataset, eval_loader = make_iwildcam_eval( + transform=eval_transform, + split=split, + batch_size=data_args.get("batch_size", 128), + root_path=root_path, + num_workers=data_args.get("num_workers", 8), + pin_mem=bool(data_args.get("pin_mem", True)), + download=bool(data_args.get("download", True)), + ) + + all_y_pred = [] + all_y_true = [] + all_metadata = [] + + with torch.no_grad(): + for batch in eval_loader: + if len(batch) == 3: + imgs, y_true, metadata = batch + else: + raise ValueError("Expected eval loader to return (x, y, metadata)") + imgs = imgs.to(device, non_blocking=True) + with torch.cuda.amp.autocast(enabled=use_autocast, dtype=torch.bfloat16): + logits = model(imgs) + preds = logits.argmax(dim=1).cpu() + all_y_pred.append(preds) + all_y_true.append(y_true.cpu()) + all_metadata.append(metadata.cpu()) + + if all_y_pred: + all_y_pred = torch.cat(all_y_pred, dim=0) + all_y_true = torch.cat(all_y_true, dim=0) + all_metadata = torch.cat(all_metadata, dim=0) if all_metadata else torch.empty(0) + + if int(all_metadata.shape[0]) != int(all_y_pred.shape[0]): + raise ValueError( + "Metadata length mismatch with predictions: " + f"{int(all_metadata.shape[0])} vs {int(all_y_pred.shape[0])}" + ) + + metrics = full_dataset.eval(all_y_pred, all_y_true, all_metadata) + + metrics_path = os.path.join(folder, f"{tag}_metrics.json") + with open(metrics_path, "w") as f: + json.dump(metrics, f, indent=2, sort_keys=True) + 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: + dist.barrier() + dist.destroy_process_group() + + +if __name__ == "__main__": + raise RuntimeError( + "Use main_eval_wilds.py to launch this script with a config file." + )