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

Class TextDataset

ddparser/parser/data_struct/data.py:72–104  ·  view source on GitHub ↗

TextDataset

Source from the content-addressed store, hash-verified

70
71
72class TextDataset(object):
73 """TextDataset"""
74 def __init__(self, corpus, fields, n_buckets=None):
75 self.corpus = corpus
76 self.fields = []
77 for field in fields:
78 if field is None:
79 continue
80 if isinstance(field, Iterable):
81 self.fields.extend(field)
82 else:
83 self.fields.append(field)
84
85 for field in self.fields:
86 setattr(self, field.name, field.transform(getattr(corpus, field.name)))
87
88 if n_buckets:
89 self.lengths = [len(i) + int(bool(field.bos)) for i in corpus]
90 self.buckets = dict(zip(*utils.kmeans(self.lengths, n_buckets)))
91
92 def __getitem__(self, index):
93 """Returns an iterator containing all fileds of a sample"""
94 for field in self.fields:
95 yield getattr(self, field.name)[index]
96
97 def __len__(self):
98 """The dataset size"""
99 return len(self.corpus)
100
101 @classmethod
102 def collate_fn(cls, batch):
103 """Return batch samples according to field"""
104 return (field for field in zip(*batch))
105
106
107class BucketsSampler(object):

Callers 6

trainFunction · 0.90
evaluateFunction · 0.90
predictFunction · 0.90
predict_queryFunction · 0.90
parseMethod · 0.90
parse_segMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected