From 6f5c6b8f26e27fe9c4a5deee0323aeca5270d95f Mon Sep 17 00:00:00 2001 From: Yann Ahlgrim Date: Mon, 25 May 2026 16:52:18 +0200 Subject: [PATCH] del folder afterwards --- src/train_supervised.py | 28 ++++++++++++++++++++++------ 1 file changed, 22 insertions(+), 6 deletions(-) diff --git a/src/train_supervised.py b/src/train_supervised.py index a356626..4e5f34b 100644 --- a/src/train_supervised.py +++ b/src/train_supervised.py @@ -1,4 +1,5 @@ import os +import shutil import sys import yaml import logging @@ -181,6 +182,8 @@ def main(args, resume_preempt=False): with open(os.path.join(folder, "params-supervised.yaml"), "w") as f: yaml.dump(args, f) + with open(os.path.join(folder, "params.yaml"), "w") as f: + yaml.dump(args, f) save_path = os.path.join(folder, f"{tag}" + "-ep{epoch}.pth.tar") latest_path = os.path.join(folder, f"{tag}-latest.pth.tar") @@ -504,13 +507,26 @@ def main(args, resume_preempt=False): } eval_wilds_main(args=eval_args) + eval_folder = eval_args.get("logging", {}).get("folder") + supervised_params = os.path.join(folder, "params-supervised.yaml") + if eval_folder and os.path.exists(supervised_params): + try: + shutil.copy2( + supervised_params, + os.path.join(eval_folder, "params-supervised.yaml"), + ) + shutil.copy2( + supervised_params, + os.path.join(eval_folder, "params.yaml"), + ) + except OSError: + logger.warning("Could not copy supervised params to eval folder") + if os.path.exists(folder): - for fname in os.listdir(folder): - if fname.endswith(".pth.tar"): - try: - os.remove(os.path.join(folder, fname)) - except OSError: - logger.warning(f"Could not remove checkpoint file: {fname}") + try: + shutil.rmtree(folder) + except OSError: + logger.warning("Could not remove supervised run folder") if __name__ == "__main__":