Download cas9/data/utils.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 2.14 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/data/utils.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/cas9/data/utils.py
-
curl -L -o utils.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/data/utils.py
2.14 kB
| # Copyright (c) Meta Platforms, Inc. and affiliates. | |
| # All rights reserved. | |
| # | |
| # This source code is licensed under the CC-by-NC license found in the | |
| # LICENSE file in the root directory of this source tree. | |
| # This implementation is adapted from https://github.com/pytorch/data/blob/main/torchdata/stateful_dataloader/sampler.py#L132 | |
| # which is released under BSD-3 license | |
| import itertools | |
| from typing import Any, Dict, Optional | |
| import numpy as np | |
| import torch | |
| from torch import Tensor | |
| from torch.utils.data import DataLoader, Dataset, Sampler | |
| def cycle_loader(dataloader: DataLoader, sampler: Sampler = None) -> Tensor: | |
| while 1: | |
| if sampler is not None: | |
| sampler.set_epoch(np.random.randint(0, 100000)) | |
| for data in dataloader: | |
| yield data | |
| class StatefulDistributedSampler(torch.utils.data.distributed.DistributedSampler): | |
| """ | |
| From: https://github.com/pytorch/data/blob/main/torchdata/stateful_dataloader/sampler.py#L132 | |
| """ | |
| _YIELDED = "yielded" | |
| def __init__( | |
| self, | |
| dataset: Dataset, | |
| num_replicas: Optional[int] = None, | |
| rank: Optional[int] = None, | |
| shuffle: bool = True, | |
| seed: int = 0, | |
| drop_last: bool = False, | |
| ) -> None: | |
| super().__init__(dataset, num_replicas, rank, shuffle, seed, drop_last) | |
| self.yielded = 0 | |
| self.next_yielded = None | |
| def __iter__(self): | |
| self.yielded = 0 | |
| if self.next_yielded is not None: | |
| self.yielded = self.next_yielded | |
| self.next_yielded = None | |
| it = super().__iter__() | |
| for idx in itertools.islice(it, self.yielded, None): | |
| self.yielded += 1 | |
| yield idx | |
| def state_dict(self) -> Dict[str, Any]: | |
| return {self._YIELDED: self.yielded} | |
| def load_state_dict(self, state_dict: Dict[str, Any]) -> None: | |
| if self._YIELDED not in state_dict: | |
| raise ValueError("Invalid state_dict") | |
| if state_dict[self._YIELDED] < 0: | |
| raise ValueError("Cannot load state_dict with negative yielded value") | |
| self.next_yielded = state_dict[self._YIELDED] | |