evaluation pipeline

This commit is contained in:
YannAhlgrim
2026-05-14 10:52:40 +02:00
parent 4a7110ebf2
commit d22a3aba90
4 changed files with 335 additions and 0 deletions
+22
View File
@@ -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
+102
View File
@@ -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()
+29
View File
@@ -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).
+182
View File
@@ -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."
)