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
+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):