From 694184cacd80a3f976be02bcc4f9d127e980e6c5 Mon Sep 17 00:00:00 2001 From: Mohammad Amanour Rahman Date: Sun, 12 Jul 2026 18:18:51 +0600 Subject: [PATCH 1/2] Fix #422: Added --fold argument and 5-fold cross-validation logic --- UNETR/BTCV/main.py | 2 +- UNETR/BTCV/utils/data_utils.py | 16 ++++++++++++---- 2 files changed, 13 insertions(+), 5 deletions(-) diff --git a/UNETR/BTCV/main.py b/UNETR/BTCV/main.py index e31b1991..fd31fa3f 100644 --- a/UNETR/BTCV/main.py +++ b/UNETR/BTCV/main.py @@ -91,7 +91,7 @@ parser.add_argument("--resume_jit", action="store_true", help="resume training from pretrained torchscript checkpoint") parser.add_argument("--smooth_dr", default=1e-6, type=float, help="constant added to dice denominator to avoid nan") parser.add_argument("--smooth_nr", default=0.0, type=float, help="constant added to dice numerator to avoid zero") - +parser.add_argument("--fold", default=0, type=int, help="fold number for 5-fold cross-validation (0-4)") def main(): args = parser.parse_args() diff --git a/UNETR/BTCV/utils/data_utils.py b/UNETR/BTCV/utils/data_utils.py index bcdd844e..df6a58a9 100755 --- a/UNETR/BTCV/utils/data_utils.py +++ b/UNETR/BTCV/utils/data_utils.py @@ -131,12 +131,20 @@ def get_loader(args): ) loader = test_loader else: - datalist = load_decathlon_datalist(datalist_json, True, "training", base_dir=data_dir) + full_datalist = load_decathlon_datalist(datalist_json, True, "training", base_dir=data_dir) + folds = data.partition_dataset(data=full_datalist, num_partitions=5, shuffle=True, seed=42) + val_files = folds[args.fold] + + train_files = [] + for i in range(5): + if i != args.fold: + train_files.extend(folds[i]) + if args.use_normal_dataset: - train_ds = data.Dataset(data=datalist, transform=train_transform) + train_ds = data.Dataset(data=train_files, transform=train_transform) else: train_ds = data.CacheDataset( - data=datalist, transform=train_transform, cache_num=24, cache_rate=1.0, num_workers=args.workers + data=train_files, transform=train_transform, cache_num=24, cache_rate=1.0, num_workers=args.workers ) train_sampler = Sampler(train_ds) if args.distributed else None train_loader = data.DataLoader( @@ -148,7 +156,7 @@ def get_loader(args): pin_memory=True, persistent_workers=True, ) - val_files = load_decathlon_datalist(datalist_json, True, "validation", base_dir=data_dir) + val_ds = data.Dataset(data=val_files, transform=val_transform) val_sampler = Sampler(val_ds, shuffle=False) if args.distributed else None val_loader = data.DataLoader( From 248c80a8c092f8058b85401d742a40056a4991cd Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 12 Jul 2026 12:21:43 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- UNETR/BTCV/utils/data_utils.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/UNETR/BTCV/utils/data_utils.py b/UNETR/BTCV/utils/data_utils.py index df6a58a9..1ce17514 100755 --- a/UNETR/BTCV/utils/data_utils.py +++ b/UNETR/BTCV/utils/data_utils.py @@ -134,12 +134,12 @@ def get_loader(args): full_datalist = load_decathlon_datalist(datalist_json, True, "training", base_dir=data_dir) folds = data.partition_dataset(data=full_datalist, num_partitions=5, shuffle=True, seed=42) val_files = folds[args.fold] - + train_files = [] for i in range(5): if i != args.fold: train_files.extend(folds[i]) - + if args.use_normal_dataset: train_ds = data.Dataset(data=train_files, transform=train_transform) else: @@ -156,7 +156,7 @@ def get_loader(args): pin_memory=True, persistent_workers=True, ) - + val_ds = data.Dataset(data=val_files, transform=val_transform) val_sampler = Sampler(val_ds, shuffle=False) if args.distributed else None val_loader = data.DataLoader(