import torch def custom_collate(batch): images = [item['image'] for item in batch] cls_labels = [item['cls_label'] for item in batch] # Some items might not have masks masks = [item['mask'] for item in batch] patient_ids = [item['patient_id'] for item in batch] # Stack only non-None masks; keep None as placeholder # or pad with dummy masks for consistent shape collated_batch = { 'image': torch.stack(images), 'patient_id': patient_ids } # --------- Handle cls_label (optional) ---------- if any(label is not None for label in cls_labels): # Use first non-None label as reference for shape for label in cls_labels: if label is not None: ref_shape = label.shape break dummy_label = torch.zeros(ref_shape, dtype=torch.long) cls_labels = [lbl if lbl is not None else dummy_label for lbl in cls_labels] collated_batch['cls_label'] = torch.stack(cls_labels) collated_batch['has_cls_label'] = torch.tensor([lbl is not None for lbl in cls_labels]) else: collated_batch['cls_label'] = None collated_batch['has_cls_label'] = None # Check if at least one sample has a mask if any(mask is not None for mask in masks): # Replace None with a dummy zero mask matching the shape # ref_shape = masks[0].shape if masks[0] is not None else masks[1].shape for mask in masks: if mask is not None: ref_shape = mask.shape break dummy = torch.zeros(ref_shape, dtype=torch.long) masks = [m if m is not None else dummy for m in masks] collated_batch['mask'] = torch.stack(masks) collated_batch['has_mask'] = torch.tensor([m is not None for m in masks]) else: collated_batch['mask'] = None collated_batch['has_mask'] = None return collated_batch