diff --git a/nnunetv2/training/dataloading/data_loader.py b/nnunetv2/training/dataloading/data_loader.py index 0c51b8272..2e214233d 100644 --- a/nnunetv2/training/dataloading/data_loader.py +++ b/nnunetv2/training/dataloading/data_loader.py @@ -15,6 +15,8 @@ from nnunetv2.utilities.plans_handling.plans_handler import PlansManager from acvl_utils.cropping_and_padding.bounding_boxes import crop_and_pad_nd +from threading import Lock +global_mutex = Lock() class nnUNetDataLoader(DataLoader): def __init__(self, @@ -198,7 +200,8 @@ def generate_train_batch(self): if self.transforms is not None: with torch.no_grad(): - with threadpool_limits(limits=1, user_api=None): + # with threadpool_limits(limits=1, user_api=None): + global_mutex.acquire() data_all = torch.from_numpy(data_all).float() seg_all = torch.from_numpy(seg_all).to(torch.int16) images = [] @@ -213,6 +216,8 @@ def generate_train_batch(self): else: seg_all = torch.stack(segs) del segs, images + global_mutex.release() + return {'data': data_all, 'target': seg_all, 'keys': selected_keys} return {'data': data_all, 'target': seg_all, 'keys': selected_keys}