| 69 | |
| 70 | |
| 71 | class DatasetIterater(object): |
| 72 | def __init__(self, batches, batch_size, device): |
| 73 | self.batch_size = batch_size |
| 74 | self.batches = batches |
| 75 | self.n_batches = len(batches) // batch_size |
| 76 | self.residue = False # 记录batch数量是否为整数 |
| 77 | if len(batches) % self.n_batches != 0: |
| 78 | self.residue = True |
| 79 | self.index = 0 |
| 80 | self.device = device |
| 81 | |
| 82 | def _to_tensor(self, datas): |
| 83 | x = torch.LongTensor([_[0] for _ in datas]).to(self.device) |
| 84 | y = torch.LongTensor([_[1] for _ in datas]).to(self.device) |
| 85 | |
| 86 | # pad前的长度(超过pad_size的设为pad_size) |
| 87 | seq_len = torch.LongTensor([_[2] for _ in datas]).to(self.device) |
| 88 | return (x, seq_len), y |
| 89 | |
| 90 | def __next__(self): |
| 91 | if self.residue and self.index == self.n_batches: |
| 92 | batches = self.batches[self.index * self.batch_size: len(self.batches)] |
| 93 | self.index += 1 |
| 94 | batches = self._to_tensor(batches) |
| 95 | return batches |
| 96 | |
| 97 | elif self.index >= self.n_batches: |
| 98 | self.index = 0 |
| 99 | raise StopIteration |
| 100 | else: |
| 101 | batches = self.batches[self.index * self.batch_size: (self.index + 1) * self.batch_size] |
| 102 | self.index += 1 |
| 103 | batches = self._to_tensor(batches) |
| 104 | return batches |
| 105 | |
| 106 | def __iter__(self): |
| 107 | return self |
| 108 | |
| 109 | def __len__(self): |
| 110 | if self.residue: |
| 111 | return self.n_batches + 1 |
| 112 | else: |
| 113 | return self.n_batches |
| 114 | |
| 115 | |
| 116 | def build_iterator(dataset, config): |