MCPcopy Create free account
hub / github.com/baidu/DDParser / TextDataLoader

Class TextDataLoader

ddparser/parser/data_struct/data.py:36–69  ·  view source on GitHub ↗

TextDataLoader

Source from the content-addressed store, hash-verified

34
35
36class TextDataLoader(object):
37 """TextDataLoader"""
38 def __init__(self, dataset, batch_sampler, collate_fn, use_data_parallel=False, use_multiprocess=True):
39 self.dataset = dataset
40 self.batch_sampler = batch_sampler
41 self.fields = self.dataset.fields
42 self.collate_fn = collate_fn
43 self.use_data_parallel = use_data_parallel
44 self.dataloader = io.DataLoader.from_generator(capacity=10, return_list=True, use_multiprocess=use_multiprocess)
45 self.dataloader.set_batch_generator(self.generator_creator())
46
47 def __call__(self):
48 """call"""
49 return self.dataloader()
50
51 def generator_creator(self):
52 """Returns a generator, each iteration returns a batch of data"""
53 def __reader():
54 for batch_sample_id in self.batch_sampler:
55 batch = []
56 raw_batch = self.collate_fn([self.dataset[sample_id] for sample_id in batch_sample_id])
57 for data, field in zip(raw_batch, self.fields):
58 if isinstance(data[0], np.ndarray):
59 data = nn.pad_sequence(data, field.pad_index)
60 elif isinstance(data[0], Iterable):
61 data = [nn.pad_sequence(f, field.pad_index) for f in zip(*data)]
62 batch.append(data)
63 yield batch
64
65 return __reader
66
67 def __len__(self):
68 """Returns the number of batches"""
69 return len(self.batch_sampler)
70
71
72class TextDataset(object):

Callers 1

batchifyFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected