auto_folder

This commit is contained in:
YannAhlgrim
2026-05-22 16:08:47 +02:00
parent a4999ab823
commit 7e25f2a57f
5 changed files with 125 additions and 11 deletions
+1 -1
View File
@@ -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
@@ -55,5 +55,5 @@ validation:
eval_every: 1
logging:
folder: experiment_logs/supervised/
write_tag: linear_probe
auto_folder: true
+24 -6
View File
@@ -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()
+31 -3
View File
@@ -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(
+68
View File
@@ -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):