MCPcopy Create free account
hub / github.com/649453932/Chinese-Text-Classification-Pytorch / DatasetIterater

Class DatasetIterater

utils.py:71–113  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

69
70
71class 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
116def build_iterator(dataset, config):

Callers 1

build_iteratorFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected