from __future__ import annotations from torch.utils.data import DataLoader, DistributedSampler, Sampler from .nifti_dataset import NiftiSegDataset from sacflow.utils.distributed import get_world_size, get_rank class DistributedEvalSamplerNoPad(Sampler): """Shard evaluation data across ranks without padding/duplication. PyTorch's DistributedSampler pads samples so every rank has equal length. That is useful for training but biases validation metrics because some cases are duplicated. This sampler uses rank::world_size indices exactly once. """ def __init__(self, dataset): self.dataset = dataset self.rank = get_rank() self.world_size = get_world_size() self.indices = list(range(self.rank, len(dataset), self.world_size)) def __iter__(self): return iter(self.indices) def __len__(self): return len(self.indices) def build_loader(cfg, split: str, training: bool, require_label: bool = False, distributed: bool | None = None): """Build a NIfTI segmentation loader. Important DDP behavior: - Training loaders use DistributedSampler when world_size > 1. - Evaluation/validation loaders default to *no* DistributedSampler. This is deliberate: training-time validation is run only on rank 0, and standalone eval usually uses one rank. Using a DistributedSampler for validation without metric all-gather biases metrics to a rank-local subset. """ data_cfg = cfg["data"] ds = NiftiSegDataset(data_cfg["manifest"], split=split, cfg=data_cfg, training=training, require_label=require_label) if distributed is None: distributed = bool(training and get_world_size() > 1) if distributed: sampler = DistributedSampler(ds, shuffle=True) if training else DistributedEvalSamplerNoPad(ds) else: sampler = None loader = DataLoader( ds, batch_size=data_cfg.get("batch_size" if training else "val_batch_size", 1), shuffle=(training and sampler is None), sampler=sampler, num_workers=cfg.get("num_workers", 4), pin_memory=cfg.get("pin_memory", True), persistent_workers=cfg.get("num_workers", 4) > 0, ) return loader