TextDataset
| 70 | |
| 71 | |
| 72 | class 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 | |
| 107 | class BucketsSampler(object): |