| 1060 | return (len(self.sampler) + self.batch_size - 1) // self.batch_size |
| 1061 | |
| 1062 | class DistributedSequentialSampler(Sampler): |
| 1063 | def __init__(self, dataset, world_size=None, rank=None): |
| 1064 | if world_size == None: |
| 1065 | world_size = get_world_size() |
| 1066 | if rank == None: |
| 1067 | rank = get_rank() |
| 1068 | self.dataset = dataset |
| 1069 | self.world_size = world_size |
| 1070 | self.rank = rank |
| 1071 | assert len(self.dataset) >= self.world_size, f'{len(self.dataset)} vs {self.world_size}' |
| 1072 | sub_num = int(math.ceil(len(self.dataset) * 1.0 / self.world_size)) |
| 1073 | self.beg = sub_num * self.rank |
| 1074 | self.end = min(self.beg+sub_num, len(self.dataset)) |
| 1075 | |
| 1076 | def __iter__(self): |
| 1077 | indices = list(range(self.beg, self.end)) |
| 1078 | return iter(indices) |
| 1079 | |
| 1080 | def __len__(self): |
| 1081 | return self.end - self.beg |
| 1082 | |
| 1083 | def simple_group_split(world_size, rank, num_groups): |
| 1084 | groups = [] |
no outgoing calls