Download sacflow/data/loader.py from sathiiii/SACFlow: direct link, hf CLI and curl.
- Browser
- Download file 2.24 kB
-
https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/data/loader.py
- Command line
-
hf download hf://sathiiii/SACFlow/sacflow/data/loader.py
-
curl -L -o loader.py https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/data/loader.py
2.24 kB
| 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 | |