| 1123 | |
| 1124 | |
| 1125 | class TestIterator: |
| 1126 | def __init__(self, iters_per_epoch, batch_size, total_iter_num=-1): |
| 1127 | self.n = iters_per_epoch |
| 1128 | self.total_n = total_iter_num |
| 1129 | self.batch_size = batch_size |
| 1130 | |
| 1131 | def __iter__(self): |
| 1132 | self.i = 0 |
| 1133 | return self |
| 1134 | |
| 1135 | def __next__(self): |
| 1136 | batch = [] |
| 1137 | # setting -1 means that no total iteration limit is set |
| 1138 | if self.i < self.n and self.total_n != 0: |
| 1139 | batch = [np.arange(0, 10, dtype=np.uint8) for _ in range(self.batch_size)] |
| 1140 | self.i += 1 |
| 1141 | self.total_n -= 1 |
| 1142 | return batch |
| 1143 | else: |
| 1144 | self.i = 0 |
| 1145 | raise StopIteration |
| 1146 | |
| 1147 | next = __next__ |
| 1148 | |
| 1149 | @property |
| 1150 | def size( |
| 1151 | self, |
| 1152 | ): |
| 1153 | return self.n * self.batch_size |
| 1154 | |
| 1155 | |
| 1156 | @nottest |
no outgoing calls